Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cb46d63f07 | ||
|
|
bbbb7ab021 |
@@ -1,23 +1,20 @@
|
||||
---
|
||||
name: reseed-performance-baseline
|
||||
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.
|
||||
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.
|
||||
---
|
||||
|
||||
# Re-seed Performance Baseline
|
||||
|
||||
## Purpose
|
||||
|
||||
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`.
|
||||
Replace or advance the rolling performance baseline for a single
|
||||
`(model_id, gpu_type)` pair in the HF dataset
|
||||
`FastVideo/performance-tracking`.
|
||||
|
||||
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`.
|
||||
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`.
|
||||
|
||||
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
|
||||
@@ -25,13 +22,11 @@ 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.
|
||||
|
||||
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`.
|
||||
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.
|
||||
|
||||
These records are intentional operator-approved baseline resets, not ordinary
|
||||
independent main-branch persistence. Mark them clearly with provenance fields
|
||||
@@ -71,8 +66,8 @@ approval, then upload reviewed accepted baseline records.
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `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. |
|
||||
| `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. |
|
||||
| `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: `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. |
|
||||
@@ -83,14 +78,10 @@ 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`, `baseline_eligible=true` records
|
||||
for the same legacy `(model_id, gpu_type)` target or the same v2 exact
|
||||
comparable identity.
|
||||
- Baseline window: last 5 `success=true` records for the same
|
||||
`(model_id, gpu_type)`.
|
||||
- Reseed count: dynamic. Upload exactly one accepted seed record per validated
|
||||
source JSON.
|
||||
|
||||
@@ -124,24 +115,12 @@ with open(source_result, encoding="utf-8") as f:
|
||||
record = json.load(f)
|
||||
```
|
||||
|
||||
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.
|
||||
Stop if any normalized record's `model_id` or `gpu_type` does not match the
|
||||
requested `model_id` and `gpu_type`.
|
||||
|
||||
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,
|
||||
@@ -215,7 +194,7 @@ export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
For legacy records, back up the sanitized model directory under `/tmp`:
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
@@ -230,16 +209,6 @@ 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
|
||||
@@ -262,9 +231,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
|
||||
### 3. Compute old baseline and candidate shift
|
||||
|
||||
Load the last 5 successful baseline records for the target.
|
||||
|
||||
For legacy targets:
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
@@ -275,28 +242,6 @@ 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,
|
||||
)
|
||||
```
|
||||
|
||||
@@ -312,8 +257,7 @@ 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 an existing last-5 median by itself, but
|
||||
it is enough to establish the first v2 baseline for a new exact identity.
|
||||
- 1 seed record usually does not move a last-5 median by itself.
|
||||
- 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
|
||||
@@ -323,10 +267,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 `<target description>`.
|
||||
> About to RE-SEED performance baseline for `<model_id>` on `<gpu_type>`.
|
||||
> This will upload `<N>` new `success=true` records to
|
||||
> `FastVideo/performance-tracking/<sanitize(model_id)>/` or the source
|
||||
> artifact's v2 model directory, one per accepted source JSON.
|
||||
> `FastVideo/performance-tracking/<sanitize(model_id)>/`, one per accepted
|
||||
> source JSON.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> Source results: `<source_results>`
|
||||
@@ -344,63 +288,25 @@ 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.
|
||||
|
||||
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
|
||||
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
|
||||
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.
|
||||
`(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.
|
||||
|
||||
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, fall back to this default
|
||||
baseline field list:
|
||||
If there are no previous HF records for the target model/GPU, fall back to this
|
||||
default baseline field list:
|
||||
|
||||
- `model_id`
|
||||
- `timestamp`
|
||||
@@ -413,22 +319,6 @@ 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.
|
||||
|
||||
@@ -444,22 +334,6 @@ 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
|
||||
@@ -482,8 +356,7 @@ Prefer uploading new accepted seed records so failed history remains visible.
|
||||
Print:
|
||||
|
||||
- Backup directory path under `/tmp`.
|
||||
- Prepared local record paths under `PERFORMANCE_RESEED_STAGING_ROOT`.
|
||||
- Prepared upload-manifest path under the identity reservation.
|
||||
- Prepared local record paths under `PERFORMANCE_TRACKING_ROOT`.
|
||||
- HF paths that will receive the new records.
|
||||
- Old rolling medians.
|
||||
- Source batch medians, source batch spread, reseed count, and candidate
|
||||
@@ -495,36 +368,22 @@ prepared records plus backup on disk.
|
||||
|
||||
### 7. Upload only the scoped records
|
||||
|
||||
For a first v2 calibration seed, use the manifest uploader after the user
|
||||
replies exactly `upload`:
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```bash
|
||||
python -c 'from fastvideo.tests.performance.seed_baseline import upload_prepared_seed_manifest; print(upload_prepared_seed_manifest("<prepared_manifest>"))'
|
||||
```python
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
|
||||
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:
|
||||
Run it once per prepared record. Each upload goes to:
|
||||
|
||||
```text
|
||||
FastVideo/performance-tracking/<sanitize(model_id)>/<record_filename>.json
|
||||
```
|
||||
|
||||
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.
|
||||
Never bulk upload the whole tracking root. Never modify another model's
|
||||
directory in the same operation.
|
||||
|
||||
### 8. Report outcome and offer cleanup
|
||||
|
||||
@@ -546,14 +405,9 @@ 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`: 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.
|
||||
- `PERFORMANCE_TRACKING_ROOT`, usually `/tmp/perf-tracking`: local synced
|
||||
mirror of `FastVideo/performance-tracking` plus the prepared local seed
|
||||
records used for scoped upload.
|
||||
- `/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.
|
||||
@@ -563,18 +417,14 @@ local state. Explain what each directory is for:
|
||||
Ask:
|
||||
|
||||
> Reseed succeeded. Do you want me to delete the local temp tracking mirror,
|
||||
> 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.
|
||||
> 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 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.
|
||||
directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
|
||||
## Failure modes and handling
|
||||
|
||||
@@ -586,34 +436,19 @@ 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 an existing last-5
|
||||
median. This warning does not block a first v2 calibration seed for an exact
|
||||
identity with no eligible baseline yet.
|
||||
clear that one or two records may not immediately move the last-5 median.
|
||||
- **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 prepared seed
|
||||
records under `/tmp/performance_reseed_prepared`, 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 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.
|
||||
@@ -624,9 +459,8 @@ remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `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/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
- `.buildkite/performance-benchmarks/tests/*.json` — fixed absolute benchmark
|
||||
@@ -639,4 +473,3 @@ 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. |
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"benchmark_version": 2,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
@@ -436,7 +436,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
|
||||
@@ -76,7 +76,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} 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"
|
||||
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"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
|
||||
+2
-3
@@ -72,7 +72,8 @@ docs/distillation/examples/
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
@@ -126,8 +127,6 @@ 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
|
||||
|
||||
@@ -67,7 +67,7 @@ MODEL_REGISTRY = {
|
||||
},
|
||||
}
|
||||
|
||||
DEFAULT_MODEL_ID = "fast-ltx23"
|
||||
DEFAULT_MODEL_ID = "fast-ltx2"
|
||||
|
||||
ACTIVE_MODEL_ID = (os.getenv("DREAMVERSE_MODEL_ID", "").strip() or DEFAULT_MODEL_ID)
|
||||
if ACTIVE_MODEL_ID not in MODEL_REGISTRY:
|
||||
@@ -76,11 +76,14 @@ if ACTIVE_MODEL_ID not in MODEL_REGISTRY:
|
||||
# Active model configuration
|
||||
MODEL_CONFIG = MODEL_REGISTRY[ACTIVE_MODEL_ID]
|
||||
|
||||
# Generation limits
|
||||
SESSION_TIMEOUT_SECONDS = 300
|
||||
|
||||
# Frame settings
|
||||
NUM_FRAMES = 121
|
||||
FRAME_HEIGHT = 1088
|
||||
FRAME_WIDTH = 1920
|
||||
NUM_INFERENCE_STEPS = 6
|
||||
NUM_INFERENCE_STEPS = 5
|
||||
JPEG_QUALITY = 100
|
||||
BATCH_SIZE = 3
|
||||
|
||||
@@ -165,9 +168,6 @@ def _optional_env(*names: str) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
# Generation limits
|
||||
SESSION_TIMEOUT_SECONDS = _env_int("DREAMVERSE_SESSION_TIMEOUT_SECONDS", 300)
|
||||
|
||||
DEVTOOLS_ENABLED = _env_bool("FASTVIDEO_ENABLE_DEVTOOLS", False)
|
||||
PROMPT_SAFETY_ENABLED = _env_bool("FASTVIDEO_ENABLE_PROMPT_SAFETY", False)
|
||||
DREAMVERSE_MAX_AUTOTUNE = _env_bool("DREAMVERSE_MAX_AUTOTUNE", True)
|
||||
|
||||
@@ -1023,11 +1023,7 @@ def get_available_gpus() -> list[int]:
|
||||
"""Get list of available GPU IDs from environment or auto-detect."""
|
||||
cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "")
|
||||
if cuda_visible:
|
||||
try:
|
||||
visible_gpu_ids = [int(x.strip()) for x in cuda_visible.split(",") if x.strip()]
|
||||
except ValueError as exc:
|
||||
raise RuntimeError("CUDA_VISIBLE_DEVICES must be a comma-separated list of integer GPU "
|
||||
f"indices (got {cuda_visible!r}); GPU UUIDs are not supported.") from exc
|
||||
visible_gpu_ids = [int(x.strip()) for x in cuda_visible.split(",") if x.strip()]
|
||||
return _limit_gpu_ids(visible_gpu_ids)
|
||||
|
||||
# Auto-detect available GPUs
|
||||
|
||||
@@ -206,9 +206,7 @@ def cli() -> None:
|
||||
args = parser.parse_args()
|
||||
|
||||
_install_heartbeat_log_filter()
|
||||
# A 15MB init image (session_init_image.MAX_SESSION_INIT_IMAGE_BYTES) is ~20MB
|
||||
# as a base64 ws message, above uvicorn's default 16MiB frame cap.
|
||||
uvicorn.run(app, host=args.host, port=args.port, ws_max_size=32 * 1024 * 1024)
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -30,11 +30,11 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from dreamverse._deps import require_dreamverse_runtime_deps
|
||||
from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP, SESSION_TIMEOUT_SECONDS
|
||||
from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP
|
||||
from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image
|
||||
from dreamverse.utils import _resolve_generation_segment_cap
|
||||
|
||||
LATENCY_MS = 200
|
||||
SESSION_TIMEOUT_SECONDS = 300
|
||||
MOCK_FRAME_WIDTH = 640
|
||||
MOCK_FRAME_HEIGHT = 352
|
||||
MOCK_FPS = 24
|
||||
@@ -334,7 +334,6 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
auto_extension_enabled = bool(init_data.get("auto_extension_enabled", False))
|
||||
loop_generation_enabled = bool(init_data.get("loop_generation_enabled", False))
|
||||
single_clip_mode = bool(init_data.get("single_clip_mode", False))
|
||||
manual_continuation_mode = bool(init_data.get("manual_continuation_mode", False))
|
||||
generation_paused = False
|
||||
|
||||
if init_type == "session_init_v2":
|
||||
@@ -345,8 +344,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
incoming_prompts = []
|
||||
|
||||
curated_prompts = [prompt.strip() for prompt in incoming_prompts if isinstance(prompt, str) and prompt.strip()]
|
||||
generation_paused = bool(not manual_continuation_mode and initial_rollout_prompt and not single_clip_mode
|
||||
and len(curated_prompts) == 0)
|
||||
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
|
||||
|
||||
try:
|
||||
session_init_image = persist_session_init_image(init_data.get("initial_image"))
|
||||
@@ -375,10 +373,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
prompt_sources_blocked = False
|
||||
pending_seed_reset = False
|
||||
pending_seed_reset_reason = ""
|
||||
pending_simple_submission: PromptSubmission | None = (PromptSubmission(
|
||||
prompt_id=str(init_data.get("initial_rollout_prompt_id") or uuid.uuid4()),
|
||||
raw_prompt=initial_rollout_prompt,
|
||||
) if manual_continuation_mode and initial_rollout_prompt else None)
|
||||
pending_simple_submission: PromptSubmission | None = None
|
||||
single_clip_waiting_for_request = False
|
||||
rollout_waiting_for_rewrite = False
|
||||
initial_rollout_waiting_for_rewrite = generation_paused
|
||||
@@ -398,26 +393,14 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
|
||||
async def send_stream_start(seed_reason: str) -> None:
|
||||
await ws_send_json({
|
||||
"type":
|
||||
"ltx2_stream_start",
|
||||
"total_segments":
|
||||
len(curated_prompts),
|
||||
"preset_id":
|
||||
preset_id,
|
||||
"stream_mode":
|
||||
"av_fmp4",
|
||||
"live_mode":
|
||||
True,
|
||||
"loop_generation_enabled":
|
||||
loop_generation_enabled,
|
||||
"loop_iteration":
|
||||
loop_iteration,
|
||||
"generation_segment_cap":
|
||||
_resolve_generation_segment_cap(
|
||||
single_clip_mode=single_clip_mode,
|
||||
cap=GENERATION_SEGMENT_CAP,
|
||||
manual_continuation_mode=manual_continuation_mode,
|
||||
),
|
||||
"type": "ltx2_stream_start",
|
||||
"total_segments": len(curated_prompts),
|
||||
"preset_id": preset_id,
|
||||
"stream_mode": "av_fmp4",
|
||||
"live_mode": True,
|
||||
"loop_generation_enabled": loop_generation_enabled,
|
||||
"loop_iteration": loop_iteration,
|
||||
"generation_segment_cap": 0,
|
||||
})
|
||||
if seed_reason == "init":
|
||||
await ws_send_json({
|
||||
@@ -510,7 +493,6 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
nonlocal auto_extension_enabled
|
||||
nonlocal loop_generation_enabled
|
||||
nonlocal single_clip_mode
|
||||
nonlocal manual_continuation_mode
|
||||
nonlocal generation_paused
|
||||
nonlocal seed_prompt_memory
|
||||
nonlocal curated_prompts
|
||||
@@ -555,7 +537,6 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
auto_extension_enabled = bool(payload.get("auto_extension_enabled", False))
|
||||
loop_generation_enabled = bool(payload.get("loop_generation_enabled", False))
|
||||
single_clip_mode = bool(payload.get("single_clip_mode", False))
|
||||
manual_continuation_mode = bool(payload.get("manual_continuation_mode", False))
|
||||
|
||||
seed_prompt_memory = list(next_curated_prompts)
|
||||
curated_prompts = list(seed_prompt_memory)
|
||||
@@ -563,14 +544,10 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
segment_idx = 0
|
||||
pending_seed_reset = False
|
||||
pending_seed_reset_reason = ""
|
||||
pending_simple_submission = (PromptSubmission(
|
||||
prompt_id=str(payload.get("initial_rollout_prompt_id") or uuid.uuid4()),
|
||||
raw_prompt=initial_rollout_prompt,
|
||||
) if manual_continuation_mode and initial_rollout_prompt else None)
|
||||
pending_simple_submission = None
|
||||
single_clip_waiting_for_request = False
|
||||
rollout_waiting_for_rewrite = False
|
||||
generation_paused = bool(not manual_continuation_mode and initial_rollout_prompt and not single_clip_mode
|
||||
and len(curated_prompts) == 0)
|
||||
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
|
||||
initial_rollout_waiting_for_rewrite = generation_paused
|
||||
rewrite_restart_pending = False
|
||||
loop_iteration = 0
|
||||
@@ -991,11 +968,10 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
project_stream_started = True
|
||||
await send_stream_start(pending_seed_reset_reason)
|
||||
pending_seed_reset_reason = ""
|
||||
|
||||
if pending_simple_submission is not None:
|
||||
submission = pending_simple_submission
|
||||
pending_simple_submission = None
|
||||
await promote_submission_to_ready(submission)
|
||||
if pending_simple_submission is not None:
|
||||
submission = pending_simple_submission
|
||||
pending_simple_submission = None
|
||||
await promote_submission_to_ready(submission)
|
||||
|
||||
if generation_paused:
|
||||
await asyncio.sleep(0.05)
|
||||
@@ -1005,8 +981,8 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
await asyncio.sleep(0.05)
|
||||
continue
|
||||
|
||||
if (not single_clip_mode and not manual_continuation_mode and not rollout_waiting_for_rewrite
|
||||
and GENERATION_SEGMENT_CAP > 0 and segment_idx >= GENERATION_SEGMENT_CAP):
|
||||
if (not single_clip_mode and not rollout_waiting_for_rewrite and GENERATION_SEGMENT_CAP > 0
|
||||
and segment_idx >= GENERATION_SEGMENT_CAP):
|
||||
rollout_waiting_for_rewrite = True
|
||||
loop_generation_enabled = False
|
||||
project_stream_started = False
|
||||
@@ -1241,9 +1217,7 @@ def cli() -> None:
|
||||
print(f"Starting mock server with {LATENCY_MS}ms latency on port {args.port}")
|
||||
|
||||
_install_heartbeat_log_filter()
|
||||
# A 15MB init image (session_init_image.MAX_SESSION_INIT_IMAGE_BYTES) is ~20MB
|
||||
# as a base64 ws message, above uvicorn's default 16MiB frame cap.
|
||||
uvicorn.run(app, host="0.0.0.0", port=args.port, ws_max_size=32 * 1024 * 1024)
|
||||
uvicorn.run(app, host="0.0.0.0", port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -313,33 +313,6 @@ def _extract_content_or_empty(response_json: dict[str, Any]) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def _find_balanced_object_end(text: str, start: int) -> int:
|
||||
"""Return the index just past the brace-balanced span opening at
|
||||
``text[start] == '{'``, honoring JSON string literals and escapes, or -1
|
||||
if the braces never balance (i.e. the object was truncated)."""
|
||||
depth = 0
|
||||
in_string = False
|
||||
escaped = False
|
||||
for i in range(start, len(text)):
|
||||
ch = text[i]
|
||||
if in_string:
|
||||
if escaped:
|
||||
escaped = False
|
||||
elif ch == "\\":
|
||||
escaped = True
|
||||
elif ch == '"':
|
||||
in_string = False
|
||||
elif ch == '"':
|
||||
in_string = True
|
||||
elif ch == "{":
|
||||
depth += 1
|
||||
elif ch == "}":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return i + 1
|
||||
return -1
|
||||
|
||||
|
||||
def _parse_json_response(content: str) -> dict[str, Any]:
|
||||
text = content.strip()
|
||||
if not text:
|
||||
@@ -357,7 +330,6 @@ def _parse_json_response(content: str) -> dict[str, Any]:
|
||||
r"```(?:json)?\s*([\s\S]*?)```",
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
last_fenced: dict[str, Any] | None = None
|
||||
for match in fence_pattern.finditer(text):
|
||||
block = match.group(1).strip()
|
||||
if not block:
|
||||
@@ -365,37 +337,21 @@ def _parse_json_response(content: str) -> dict[str, Any]:
|
||||
try:
|
||||
parsed = json.loads(block)
|
||||
if isinstance(parsed, dict):
|
||||
last_fenced = parsed
|
||||
return parsed
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if last_fenced is not None:
|
||||
return last_fenced
|
||||
|
||||
# Scan for all decodable JSON objects and return the last — chain-of-thought
|
||||
# models emit draft JSON mid-reasoning; the final answer is always last.
|
||||
# Fall back to scanning for the first decodable JSON object in free-form text.
|
||||
decoder = json.JSONDecoder()
|
||||
last_parsed: dict[str, Any] | None = None
|
||||
pos = 0
|
||||
while (idx := text.find("{", pos)) != -1:
|
||||
for idx, char in enumerate(text):
|
||||
if char != "{":
|
||||
continue
|
||||
try:
|
||||
parsed, consumed = decoder.raw_decode(text[idx:])
|
||||
parsed, _ = decoder.raw_decode(text[idx:])
|
||||
except json.JSONDecodeError:
|
||||
# Skip the whole failed object rather than rescanning inside it:
|
||||
# fragments nested in a malformed or truncated (finish_reason=
|
||||
# length) object must not override an earlier complete object.
|
||||
span_end = _find_balanced_object_end(text, idx)
|
||||
if span_end == -1:
|
||||
break
|
||||
pos = span_end
|
||||
continue
|
||||
# Skip past the consumed span so nested braces inside a decoded
|
||||
# object are not re-parsed as standalone objects.
|
||||
pos = idx + consumed
|
||||
if isinstance(parsed, dict):
|
||||
last_parsed = parsed
|
||||
|
||||
if last_parsed is not None:
|
||||
return last_parsed
|
||||
return parsed
|
||||
|
||||
raise ValueError("No JSON object found in assistant response.")
|
||||
|
||||
@@ -425,7 +381,7 @@ def _format_locked_segments(locked_segments: list[str]) -> str:
|
||||
|
||||
class PromptEnhancer:
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self):
|
||||
self.provider = PROMPT_PROVIDER
|
||||
self.provider_label = _resolve_provider_label(PROMPT_PROVIDER)
|
||||
self.api_key = PROMPT_API_KEY
|
||||
@@ -1461,12 +1417,15 @@ class PromptEnhancer:
|
||||
locked_text = _format_locked_segments(locked_segments_clean)
|
||||
request_system_prompt = self.enhance_system_prompt
|
||||
user_payload = {
|
||||
"request": ("<locked_segments>\n"
|
||||
f"{locked_text}\n"
|
||||
"</locked_segments>\n\n"
|
||||
f"<conditioning_prompt>{cleaned}</conditioning_prompt>\n\n"
|
||||
f"Write exactly one new segment ({next_segment_key}) "
|
||||
"continuing from the locked segments."),
|
||||
"request": (
|
||||
"<locked_segments>\n"
|
||||
f"{locked_text}\n"
|
||||
"</locked_segments>\n\n"
|
||||
f"<conditioning_prompt>{cleaned}</conditioning_prompt>\n\n"
|
||||
f"Write exactly one new segment ({next_segment_key}) "
|
||||
"continuing from the locked segments. "
|
||||
'Respond with valid JSON only as {"next_prompt": "..."}.' # noqa: E501
|
||||
),
|
||||
}
|
||||
|
||||
t0 = time.perf_counter()
|
||||
@@ -1498,7 +1457,6 @@ class PromptEnhancer:
|
||||
body=request_body,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
_enhance_print("INFO", f"raw_response: {response_content}")
|
||||
if is_single_clip_mode:
|
||||
prompt = self._extract_single_clip_prompt(response_content)
|
||||
else:
|
||||
@@ -1581,7 +1539,7 @@ class PromptEnhancer:
|
||||
f"Write exactly one new segment ({next_segment_key}) "
|
||||
"that continues linearly from the locked segments. "
|
||||
"Infer the next narrative beat from this history. "
|
||||
'Respond with valid JSON only: {"next_prompt": "<your segment description here>"}.' # noqa: E501
|
||||
'Respond with valid JSON only as {"next_prompt": "..."}.' # noqa: E501
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@@ -20,7 +20,6 @@ from __future__ import annotations
|
||||
# mypy: ignore-errors
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
@@ -53,14 +52,6 @@ if TYPE_CHECKING:
|
||||
from dreamverse.prompt_enhancer import PromptEnhancer
|
||||
from dreamverse.prompt_safety import PromptSafetyFilter
|
||||
|
||||
# Optional append-only log of every generated segment prompt; unset disables it.
|
||||
SEGMENT_PROMPT_LOG_PATH = os.environ.get("DREAMVERSE_SEGMENT_PROMPT_LOG", "")
|
||||
|
||||
|
||||
def _append_segment_prompt_log(path: str, text: str) -> None:
|
||||
with open(path, "a") as f:
|
||||
f.write(text)
|
||||
|
||||
|
||||
class SessionController:
|
||||
"""Runs one WebSocket session from accept() through disconnect."""
|
||||
@@ -207,7 +198,6 @@ class SessionController:
|
||||
auto_extension_enabled = bool(init_data.get("auto_extension_enabled", False))
|
||||
loop_generation_enabled = bool(init_data.get("loop_generation_enabled", False))
|
||||
single_clip_mode = bool(init_data.get("single_clip_mode", False))
|
||||
manual_continuation_mode = bool(init_data.get("manual_continuation_mode", False))
|
||||
rewrite_model = self.prompt_enhancer.resolve_rewrite_model(init_data.get("rewrite_model"))
|
||||
rewrite_system_prompt_override = str(init_data.get("rewrite_window_system_prompt") or "").strip()
|
||||
rewrite_user_system_prompt_override = str(init_data.get("rewrite_user_system_prompt") or "").strip()
|
||||
@@ -292,9 +282,6 @@ class SessionController:
|
||||
# Session queues and mutable state.
|
||||
raw_prompt_queue: asyncio.Queue[PromptSubmission] = asyncio.Queue()
|
||||
ready_prompt_queue: asyncio.Queue[ReadyPrompt] = asyncio.Queue()
|
||||
# Submissions dequeued by prompt_worker_loop but not yet resolved; while
|
||||
# non-zero the prompt sources are busy, not drained.
|
||||
prompt_enhancement_inflight = 0
|
||||
|
||||
curated_idx = 0
|
||||
segment_idx = 0
|
||||
@@ -304,20 +291,13 @@ class SessionController:
|
||||
generation_cap_blocked = False
|
||||
auto_extension_blocked_segment_idx: int | None = None
|
||||
prompt_sources_drained_logged = False
|
||||
generation_paused = bool(not manual_continuation_mode and initial_rollout_prompt and not single_clip_mode
|
||||
and len(curated_prompts) == 0)
|
||||
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
|
||||
pending_seed_reset = False
|
||||
pending_seed_reset_reason = ""
|
||||
pending_reset_conditioning = False
|
||||
loop_iteration = 0 if generation_paused else 1
|
||||
force_curated_restart_segment = False
|
||||
# The frontend records the opening scene under this id; reuse it so
|
||||
# prompt lifecycle events for the opening prompt reach that record.
|
||||
pending_simple_prompt_submission: PromptSubmission | None = (PromptSubmission(
|
||||
prompt_id=str(init_data.get("initial_rollout_prompt_id") or uuid.uuid4()),
|
||||
raw_prompt=initial_rollout_prompt,
|
||||
created_at_s=time.time(),
|
||||
) if manual_continuation_mode and initial_rollout_prompt else None)
|
||||
pending_simple_prompt_submission: PromptSubmission | None = None
|
||||
single_clip_waiting_for_request = False
|
||||
rollout_waiting_for_rewrite = False
|
||||
initial_rollout_waiting_for_rewrite = generation_paused
|
||||
@@ -326,7 +306,6 @@ class SessionController:
|
||||
project_active = True
|
||||
project_stream_started = False
|
||||
pending_project_end = False
|
||||
segment_prompt_log_warned = False
|
||||
|
||||
def replace_session_init_image(initial_image_payload: object) -> None:
|
||||
nonlocal session_init_image
|
||||
@@ -445,7 +424,6 @@ class SessionController:
|
||||
nonlocal auto_extension_enabled
|
||||
nonlocal loop_generation_enabled
|
||||
nonlocal single_clip_mode
|
||||
nonlocal manual_continuation_mode
|
||||
nonlocal generation_paused
|
||||
nonlocal curated_prompts
|
||||
nonlocal seed_prompt_memory
|
||||
@@ -480,7 +458,6 @@ class SessionController:
|
||||
next_auto_extension_enabled = bool(payload.get("auto_extension_enabled", False))
|
||||
next_loop_generation_enabled = bool(payload.get("loop_generation_enabled", False))
|
||||
next_single_clip_mode = bool(payload.get("single_clip_mode", False))
|
||||
next_manual_continuation_mode = bool(payload.get("manual_continuation_mode", False))
|
||||
|
||||
next_preset_id = str(payload.get("preset_id") or "").strip()
|
||||
if next_preset_id:
|
||||
@@ -533,7 +510,6 @@ class SessionController:
|
||||
auto_extension_enabled = next_auto_extension_enabled
|
||||
loop_generation_enabled = next_loop_generation_enabled
|
||||
single_clip_mode = next_single_clip_mode
|
||||
manual_continuation_mode = next_manual_continuation_mode
|
||||
rewrite_model = next_rewrite_model
|
||||
rewrite_system_prompt_override = (next_rewrite_system_prompt_override)
|
||||
rewrite_user_system_prompt_override = (next_rewrite_user_system_prompt_override)
|
||||
@@ -554,15 +530,10 @@ class SessionController:
|
||||
generation_cap_blocked = False
|
||||
auto_extension_blocked_segment_idx = None
|
||||
prompt_sources_drained_logged = False
|
||||
pending_simple_prompt_submission = (PromptSubmission(
|
||||
prompt_id=str(payload.get("initial_rollout_prompt_id") or uuid.uuid4()),
|
||||
raw_prompt=initial_rollout_prompt,
|
||||
created_at_s=time.time(),
|
||||
) if manual_continuation_mode and initial_rollout_prompt else None)
|
||||
pending_simple_prompt_submission = None
|
||||
single_clip_waiting_for_request = False
|
||||
rollout_waiting_for_rewrite = False
|
||||
generation_paused = bool(not manual_continuation_mode and initial_rollout_prompt
|
||||
and not single_clip_mode and len(curated_prompts) == 0)
|
||||
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
|
||||
initial_rollout_waiting_for_rewrite = generation_paused
|
||||
rewrite_restart_pending = False
|
||||
loop_iteration = 0
|
||||
@@ -971,7 +942,6 @@ class SessionController:
|
||||
_resolve_generation_segment_cap(
|
||||
single_clip_mode=single_clip_mode,
|
||||
cap=GENERATION_SEGMENT_CAP,
|
||||
manual_continuation_mode=manual_continuation_mode,
|
||||
),
|
||||
})
|
||||
continue
|
||||
@@ -1035,9 +1005,11 @@ class SessionController:
|
||||
continue
|
||||
|
||||
async def prompt_worker_loop():
|
||||
nonlocal prompt_enhancement_inflight
|
||||
|
||||
async def process_submission(submission: PromptSubmission) -> None:
|
||||
while not stop_event.is_set():
|
||||
try:
|
||||
submission = await asyncio.wait_for(raw_prompt_queue.get(), timeout=0.1)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
_main_print("INFO", f"Received user prompt for enhancement: {submission.raw_prompt}")
|
||||
prompt_id = submission.prompt_id
|
||||
raw_prompt = submission.raw_prompt
|
||||
@@ -1060,9 +1032,8 @@ class SessionController:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"message": blocked_raw_prompt_error,
|
||||
"prompt_id": prompt_id,
|
||||
})
|
||||
return
|
||||
continue
|
||||
await log_event(
|
||||
"enhance_request",
|
||||
{
|
||||
@@ -1122,9 +1093,8 @@ class SessionController:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"message": blocked_final_prompt_error,
|
||||
"prompt_id": prompt_id,
|
||||
})
|
||||
return
|
||||
continue
|
||||
if result.fallback_used or not final_prompt:
|
||||
source = "user_enhancement_failed"
|
||||
_main_print(
|
||||
@@ -1143,7 +1113,7 @@ class SessionController:
|
||||
})
|
||||
# Enhancement is strict JSON-only; do not enqueue raw
|
||||
# prompt when enhancement fails.
|
||||
return
|
||||
continue
|
||||
else:
|
||||
source = "user_enhanced"
|
||||
await ws_send_json({
|
||||
@@ -1160,7 +1130,6 @@ class SessionController:
|
||||
source=source,
|
||||
fallback_used=result.fallback_used,
|
||||
loop_iteration=loop_iteration,
|
||||
raw_prompt=raw_prompt,
|
||||
))
|
||||
else:
|
||||
await ready_prompt_queue.put(
|
||||
@@ -1170,7 +1139,6 @@ class SessionController:
|
||||
source="user_raw",
|
||||
fallback_used=False,
|
||||
loop_iteration=loop_iteration,
|
||||
raw_prompt=raw_prompt,
|
||||
))
|
||||
await ws_send_json({
|
||||
"type": "prompt_ready",
|
||||
@@ -1180,21 +1148,6 @@ class SessionController:
|
||||
"latency_ms": 0.0,
|
||||
})
|
||||
|
||||
while not stop_event.is_set():
|
||||
try:
|
||||
submission = raw_prompt_queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
await asyncio.sleep(0.1)
|
||||
continue
|
||||
# Dequeue and increment without an await in between so the
|
||||
# generation loop never sees an empty queue with zero in flight
|
||||
# while this submission is still being enhanced.
|
||||
prompt_enhancement_inflight += 1
|
||||
try:
|
||||
await process_submission(submission)
|
||||
finally:
|
||||
prompt_enhancement_inflight -= 1
|
||||
|
||||
def queue_snapshot() -> dict[str, object]:
|
||||
return {
|
||||
"user_ready": ready_prompt_queue.qsize(),
|
||||
@@ -1347,7 +1300,6 @@ class SessionController:
|
||||
_resolve_generation_segment_cap(
|
||||
single_clip_mode=single_clip_mode,
|
||||
cap=GENERATION_SEGMENT_CAP,
|
||||
manual_continuation_mode=manual_continuation_mode,
|
||||
),
|
||||
})
|
||||
await ws_send_json({
|
||||
@@ -1407,7 +1359,6 @@ class SessionController:
|
||||
_resolve_generation_segment_cap(
|
||||
single_clip_mode=single_clip_mode,
|
||||
cap=GENERATION_SEGMENT_CAP,
|
||||
manual_continuation_mode=manual_continuation_mode,
|
||||
),
|
||||
})
|
||||
if nonlocal_reason == "loop_restart":
|
||||
@@ -1437,9 +1388,8 @@ class SessionController:
|
||||
await raw_prompt_queue.put(pending_simple_prompt_submission)
|
||||
pending_simple_prompt_submission = None
|
||||
|
||||
if (not single_clip_mode and not manual_continuation_mode and not generation_cap_blocked
|
||||
and not rollout_waiting_for_rewrite and GENERATION_SEGMENT_CAP > 0
|
||||
and generated_segment_count >= GENERATION_SEGMENT_CAP):
|
||||
if (not single_clip_mode and not generation_cap_blocked and not rollout_waiting_for_rewrite
|
||||
and GENERATION_SEGMENT_CAP > 0 and generated_segment_count >= GENERATION_SEGMENT_CAP):
|
||||
loop_generation_enabled = False
|
||||
rollout_waiting_for_rewrite = True
|
||||
_main_print(
|
||||
@@ -1592,10 +1542,7 @@ class SessionController:
|
||||
if single_clip_mode:
|
||||
await asyncio.sleep(PROMPT_AUTO_SLEEP_MS / 1000.0)
|
||||
continue
|
||||
# A raw submission still queued or being enhanced will produce a
|
||||
# ready prompt shortly; that is not a drained/blocked state.
|
||||
enhancement_pending = (raw_prompt_queue.qsize() > 0 or prompt_enhancement_inflight > 0)
|
||||
if not prompt_sources_drained_logged and not enhancement_pending:
|
||||
if not prompt_sources_drained_logged:
|
||||
snapshot = queue_snapshot()
|
||||
_main_print(
|
||||
"WARN",
|
||||
@@ -1634,24 +1581,6 @@ class SessionController:
|
||||
total_segments_hint = max(segment_idx, len(curated_prompts))
|
||||
prompt = selected.prompt
|
||||
locked_segment_prompts.append(prompt)
|
||||
if SEGMENT_PROMPT_LOG_PATH:
|
||||
_ts = time.strftime("%Y-%m-%d %H:%M:%S")
|
||||
_lines = [
|
||||
f"\n=== Segment {segment_idx} [{_ts}] source={selected.source} client={client_id[:8]} ===",
|
||||
]
|
||||
if selected.raw_prompt and selected.raw_prompt != prompt:
|
||||
_lines.append(f"User: {selected.raw_prompt}")
|
||||
_lines.append(f"Rewritten: {prompt}")
|
||||
try:
|
||||
await asyncio.to_thread(_append_segment_prompt_log, SEGMENT_PROMPT_LOG_PATH,
|
||||
"\n".join(_lines) + "\n")
|
||||
except Exception as exc:
|
||||
if not segment_prompt_log_warned:
|
||||
segment_prompt_log_warned = True
|
||||
_main_print(
|
||||
"WARN",
|
||||
f"Failed to write segment prompt log {SEGMENT_PROMPT_LOG_PATH}: {exc}",
|
||||
)
|
||||
if (auto_extension_blocked_segment_idx is not None
|
||||
and auto_extension_blocked_segment_idx <= segment_idx):
|
||||
auto_extension_blocked_segment_idx = None
|
||||
|
||||
@@ -19,4 +19,3 @@ class ReadyPrompt:
|
||||
fallback_used: bool = False
|
||||
seed_prompt_index: int | None = None
|
||||
loop_iteration: int | None = None
|
||||
raw_prompt: str | None = None
|
||||
|
||||
@@ -150,22 +150,12 @@ def test_config_enables_prompt_safety_when_requested(monkeypatch):
|
||||
|
||||
def test_config_uses_five_minute_session_timeout(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.delenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 300
|
||||
|
||||
|
||||
def test_config_session_timeout_env_override(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.setenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", "1800")
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 1800
|
||||
|
||||
|
||||
def test_config_rejects_invalid_prompt_provider(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_PROMPT_PROVIDER", "unsupported")
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
|
||||
@@ -75,13 +75,12 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
|
||||
calls: list[dict[str, object]] = []
|
||||
uvicorn_stub = types.ModuleType("uvicorn")
|
||||
|
||||
def run(app, host: str, port: int, **kwargs) -> None:
|
||||
def run(app, host: str, port: int) -> None:
|
||||
calls.append(
|
||||
{
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -105,7 +104,6 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
"ws_max_size": 32 * 1024 * 1024,
|
||||
}
|
||||
]
|
||||
|
||||
@@ -123,7 +121,6 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
"ws_max_size": 32 * 1024 * 1024,
|
||||
}
|
||||
]
|
||||
|
||||
@@ -150,7 +147,6 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
"ws_max_size": 32 * 1024 * 1024,
|
||||
}
|
||||
]
|
||||
|
||||
@@ -170,7 +166,6 @@ def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
"ws_max_size": 32 * 1024 * 1024,
|
||||
}
|
||||
]
|
||||
assert mock_server.LATENCY_MS == 321
|
||||
|
||||
@@ -290,152 +290,6 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
|
||||
|
||||
def test_mock_server_manual_mode_streams_initial_prompt_without_rewrite_or_cap():
|
||||
old_segment_bytes = mock_server.MOCK_SEGMENT_BYTES
|
||||
old_latency_ms = mock_server.LATENCY_MS
|
||||
old_generation_segment_cap = mock_server.GENERATION_SEGMENT_CAP
|
||||
try:
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
mock_server.GENERATION_SEGMENT_CAP = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A drone skims a neon canyon",
|
||||
"initial_rollout_prompt_id": "steer-prompt-1",
|
||||
"manual_continuation_mode": True,
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.08,
|
||||
{
|
||||
"type": "append_prompt",
|
||||
"prompt": "The drone dives toward the river",
|
||||
"prompt_id": "steer-prompt-2",
|
||||
},
|
||||
),
|
||||
(0.30, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
message_types = [payload["type"] for payload in ws.sent_json]
|
||||
assert "rewrite_seed_prompts_started" not in message_types
|
||||
assert "rewrite_seed_prompts_complete" not in message_types
|
||||
assert "ltx2_stream_start" in message_types
|
||||
# cap=1 must not stop a manual-mode session after the first segment
|
||||
assert "ltx2_stream_complete" not in message_types
|
||||
|
||||
prompt_ready_events = [
|
||||
payload for payload in ws.sent_json if payload["type"] == "prompt_ready"
|
||||
]
|
||||
assert [payload["prompt_id"] for payload in prompt_ready_events] == [
|
||||
"steer-prompt-1",
|
||||
"steer-prompt-2",
|
||||
]
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
assert [payload["prompt"] for payload in segment_start_events] == [
|
||||
"A drone skims a neon canyon",
|
||||
"The drone dives toward the river",
|
||||
]
|
||||
segment_source_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "segment_prompt_source"
|
||||
]
|
||||
assert [payload["prompt_id"] for payload in segment_source_events] == [
|
||||
"steer-prompt-1",
|
||||
"steer-prompt-2",
|
||||
]
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
mock_server.GENERATION_SEGMENT_CAP = old_generation_segment_cap
|
||||
|
||||
|
||||
def test_mock_server_project_init_manual_mode_streams_initial_prompt():
|
||||
old_segment_bytes = mock_server.MOCK_SEGMENT_BYTES
|
||||
old_latency_ms = mock_server.LATENCY_MS
|
||||
try:
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.05, {"type": "end_project_keep_session"}),
|
||||
(
|
||||
0.15,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A drone skims a neon canyon",
|
||||
"initial_rollout_prompt_id": "steer-prompt-1",
|
||||
"manual_continuation_mode": True,
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.45, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
message_types = [payload["type"] for payload in ws.sent_json]
|
||||
assert "project_idle" in message_types
|
||||
project_idle_index = message_types.index("project_idle")
|
||||
# manual-mode restart must not run the rewrite rollout
|
||||
assert "rewrite_seed_prompts_started" not in message_types[project_idle_index:]
|
||||
assert "ltx2_stream_start" in message_types[project_idle_index:]
|
||||
|
||||
prompt_ready_events = [
|
||||
payload for payload in ws.sent_json if payload["type"] == "prompt_ready"
|
||||
]
|
||||
assert [payload["prompt_id"] for payload in prompt_ready_events] == [
|
||||
"steer-prompt-1",
|
||||
]
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
assert segment_start_events[-1]["prompt"] == "A drone skims a neon canyon"
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
|
||||
|
||||
def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
old_segment_bytes = mock_server.MOCK_SEGMENT_BYTES
|
||||
old_latency_ms = mock_server.LATENCY_MS
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
import re
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
@@ -206,59 +205,6 @@ def test_parse_json_response_extracts_first_embedded_object():
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
def test_parse_json_response_returns_outer_object_not_nested_value():
|
||||
parsed = _parse_json_response(
|
||||
"Final answer: {\"next_prompt\": \"a scene\", \"style\": {\"mood\": \"noir\"}} done."
|
||||
)
|
||||
assert parsed == {"next_prompt": "a scene", "style": {"mood": "noir"}}
|
||||
|
||||
|
||||
def test_parse_json_response_returns_last_of_multiple_objects():
|
||||
parsed = _parse_json_response(
|
||||
"Draft: {\"next_prompt\": \"draft\"}\nRefined: {\"next_prompt\": \"final\"}"
|
||||
)
|
||||
assert parsed == {"next_prompt": "final"}
|
||||
|
||||
|
||||
def test_parse_json_response_ignores_fragments_of_truncated_trailing_object():
|
||||
# finish_reason=length cut the refined object short; the complete draft must
|
||||
# win over a nested fragment of the truncated object.
|
||||
parsed = _parse_json_response(
|
||||
'{"next_prompt": "draft"} refined: {"next_prompt": "final", "style": {"mood": "noir"}'
|
||||
)
|
||||
assert parsed == {"next_prompt": "draft"}
|
||||
|
||||
|
||||
def test_parse_json_response_ignores_fragments_of_mid_string_truncated_object():
|
||||
# Unterminated-string truncation reports the error at the opening quote,
|
||||
# not end-of-text; nested fragments still must not win over the draft.
|
||||
parsed = _parse_json_response(
|
||||
'{"next_prompt": "draft"} refined: {"style": {"mood": "noir"}, "next_prompt": "cut off'
|
||||
)
|
||||
assert parsed == {"next_prompt": "draft"}
|
||||
|
||||
|
||||
def test_parse_json_response_ignores_fragments_of_malformed_object_with_trailing_prose():
|
||||
parsed = _parse_json_response(
|
||||
'{"next_prompt": "draft"} {"final": {"mood": "noir"}, "x": 1 and then some prose'
|
||||
)
|
||||
assert parsed == {"next_prompt": "draft"}
|
||||
|
||||
|
||||
def test_parse_json_response_raises_when_only_object_is_truncated():
|
||||
with pytest.raises(ValueError):
|
||||
_parse_json_response('{"style": {"mood": "noir"}, "next_prompt": "cut off')
|
||||
|
||||
|
||||
def test_parse_json_response_returns_outer_rollout_dict():
|
||||
parsed = _parse_json_response(
|
||||
"{\"rollout\": {\"segment_prompts\": [{\"prompt\": \"a\"}, {\"prompt\": \"b\"}]}}"
|
||||
)
|
||||
assert parsed == {
|
||||
"rollout": {"segment_prompts": [{"prompt": "a"}, {"prompt": "b"}]}
|
||||
}
|
||||
|
||||
|
||||
def test_load_prompt_required_falls_back_to_default_path(tmp_path):
|
||||
fallback_path = tmp_path / "next_segment_system_prompt.md"
|
||||
fallback_path.write_text("fallback prompt\n", encoding="utf-8")
|
||||
|
||||
@@ -15,6 +15,5 @@ def _utc_now_iso() -> str:
|
||||
PROMPT_EXTENSION_FAILURE_USER_MESSAGE = ("Prompt extension failed for this request.")
|
||||
|
||||
|
||||
def _resolve_generation_segment_cap(*, single_clip_mode: bool, cap: int, manual_continuation_mode: bool = False) -> int:
|
||||
# Steering (manual continuation) lets the user keep going indefinitely, like single-clip mode.
|
||||
return 0 if (single_clip_mode or manual_continuation_mode) else cap
|
||||
def _resolve_generation_segment_cap(*, single_clip_mode: bool, cap: int) -> int:
|
||||
return 0 if single_clip_mode else cap
|
||||
|
||||
@@ -489,7 +489,7 @@ class VideoGenerationWorker:
|
||||
num_inference_steps=NUM_INFERENCE_STEPS,
|
||||
guidance_scale=1.0,
|
||||
seed=10,
|
||||
ltx2_image_crf=(33.0 if image_path and segment_idx == 1 else 0.0),
|
||||
ltx2_image_crf=0.0,
|
||||
image_path=image_path if segment_idx == 1 else None,
|
||||
return_continuation_state=False,
|
||||
)
|
||||
|
||||
@@ -29,9 +29,6 @@ test = [
|
||||
dreamverse-server = "dreamverse.server_entry:cli"
|
||||
dreamverse-mock-server = "dreamverse.mock_server:cli"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
include = ["dreamverse*"]
|
||||
|
||||
[tool.uv]
|
||||
package = false
|
||||
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# launch-dreamverse.sh — launch dreamverse-server on a compute node.
|
||||
#
|
||||
# Usage (from repo root):
|
||||
# bash apps/dreamverse/scripts/launch-dreamverse.sh # GPUs 0-3, SP_SIZE=4
|
||||
# CUDA_VISIBLE_DEVICES=0 bash apps/dreamverse/scripts/launch-dreamverse.sh # single GPU, SP_SIZE=1
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
CONDA_PREFIX="$HOME/miniconda3/envs/dreamverse"
|
||||
CUDA_RT_DIR="$CONDA_PREFIX/lib/python3.11/site-packages/nvidia/cuda_runtime/lib"
|
||||
GXX="$CONDA_PREFIX/bin/aarch64-conda-linux-gnu-g++"
|
||||
|
||||
|
||||
export CUDA_HOME="$CONDA_PREFIX"
|
||||
export FASTVIDEO_ENABLE_STARTUP_WARMUP=true
|
||||
export FASTVIDEO_ENABLE_PROMPT_SAFETY=false
|
||||
export DREAMVERSE_MAX_AUTOTUNE=true
|
||||
export LTX2_USE_DISTILLED_SIGMAS=0
|
||||
export LTX2_VIDEO_CONDITIONING_NUM_FRAMES=1
|
||||
export AUDIO_CONDITIONING_NUM_FRAMES=41
|
||||
export DREAMVERSE_SESSION_TIMEOUT_SECONDS="${DREAMVERSE_SESSION_TIMEOUT_SECONDS:-1800}"
|
||||
# GB200 max-autotune warmup compiles can run for hours; keep the watchdog generous here
|
||||
export FASTVIDEO_STARTUP_WARMUP_TIMEOUT_SECONDS="${FASTVIDEO_STARTUP_WARMUP_TIMEOUT_SECONDS:-24000}"
|
||||
export CEREBRAS_API_KEY="${CEREBRAS_API_KEY:-}" # set this in your env or ~/.env
|
||||
export FASTVIDEO_PROMPT_CEREBRAS_MODEL="gpt-oss-120b"
|
||||
export TORCHINDUCTOR_CACHE_DIR="$HOME/.cache/torchinductor"
|
||||
export TRITON_CACHE_DIR="$HOME/.triton/cache"
|
||||
export TORCH_CUDA_ARCH_LIST="10.0a"
|
||||
|
||||
# Compiler env (needed for flashinfer JIT compilation at server startup)
|
||||
export CXX="$CONDA_PREFIX/compiler_compat/g++"
|
||||
export CC="$CONDA_PREFIX/compiler_compat/gcc"
|
||||
export CUDAHOSTCXX="$GXX"
|
||||
export NVCC_PREPEND_FLAGS="-ccbin $GXX -allow-unsupported-compiler"
|
||||
|
||||
# Link against libcudart.so.12 at JIT compile time; stubs for libcuda.so
|
||||
# cuda-compat has libcudart.so -> libcudart.so.12 (linker needs unversioned name)
|
||||
export LIBRARY_PATH="$CONDA_PREFIX/lib/cuda-compat:$CONDA_PREFIX/lib/stubs"
|
||||
|
||||
# Only libcudart.so.12 at runtime — prevents cuDNN from seeing .so.13
|
||||
export LD_LIBRARY_PATH="$CUDA_RT_DIR"
|
||||
|
||||
CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES:-0,1,2,3}"
|
||||
export FASTVIDEO_GPU_COUNT="${FASTVIDEO_GPU_COUNT:-all}"
|
||||
# Default SP size to the usable GPU count so single-GPU invocations work.
|
||||
# A numeric FASTVIDEO_GPU_COUNT caps the pool below the visible count, and an
|
||||
# SP size above the pool size fails GPUPool startup with "Not enough GPUs".
|
||||
IFS=',' read -ra _VISIBLE_GPUS <<< "$CUDA_VISIBLE_DEVICES"
|
||||
# Count only non-empty tokens, matching gpu_pool.get_available_gpus (e.g. ",0,1" is 2 GPUs).
|
||||
_USABLE_GPU_COUNT=0
|
||||
for _gpu in "${_VISIBLE_GPUS[@]}"; do
|
||||
[[ -n "${_gpu//[[:space:]]/}" ]] && _USABLE_GPU_COUNT=$((_USABLE_GPU_COUNT + 1))
|
||||
done
|
||||
if [[ "$FASTVIDEO_GPU_COUNT" =~ ^[0-9]+$ ]] && (( FASTVIDEO_GPU_COUNT < _USABLE_GPU_COUNT )); then
|
||||
_USABLE_GPU_COUNT="$FASTVIDEO_GPU_COUNT"
|
||||
fi
|
||||
export DREAMVERSE_SP_SIZE="${DREAMVERSE_SP_SIZE:-$_USABLE_GPU_COUNT}"
|
||||
PORT="${DREAMVERSE_PORT:-8009}"
|
||||
|
||||
FFMPEG_ENV="$(dirname "$0")/ffmpeg-env.sh"
|
||||
# shellcheck source=ffmpeg-env.sh
|
||||
[[ -f "$FFMPEG_ENV" ]] && source "$FFMPEG_ENV"
|
||||
|
||||
echo "==> Launching dreamverse-server on GPU $CUDA_VISIBLE_DEVICES port $PORT"
|
||||
CUDA_VISIBLE_DEVICES="$CUDA_VISIBLE_DEVICES" \
|
||||
"$CONDA_PREFIX/bin/dreamverse-server" --host 0.0.0.0 --port "$PORT"
|
||||
@@ -1,30 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# launch-frontend.sh — start the Dreamverse Next.js dev server and ngrok tunnel.
|
||||
#
|
||||
# Usage (from repo root):
|
||||
# bash apps/dreamverse/scripts/launch-frontend.sh
|
||||
#
|
||||
# Override backend or ngrok URL via env:
|
||||
# BACKEND_HOST=1.2.3.4 bash apps/dreamverse/scripts/launch-frontend.sh
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
CONDA_PREFIX="$HOME/miniconda3/envs/dreamverse"
|
||||
BACKEND_HOST="${BACKEND_HOST:-10.244.18.228}"
|
||||
BACKEND_PORT="${BACKEND_PORT:-8009}"
|
||||
NGROK_URL="${NGROK_URL:-ltx23.ngrok.app}"
|
||||
WEB_DIR="$(git rev-parse --show-toplevel)/apps/dreamverse/web"
|
||||
|
||||
cleanup() {
|
||||
echo "==> Shutting down..."
|
||||
kill "$FRONTEND_PID" 2>/dev/null || true
|
||||
}
|
||||
trap cleanup EXIT
|
||||
|
||||
echo "==> Starting frontend (backend: $BACKEND_HOST:$BACKEND_PORT)"
|
||||
BACKEND_HOST="$BACKEND_HOST" BACKEND_PORT="$BACKEND_PORT" \
|
||||
npm run --prefix "$WEB_DIR" dev &
|
||||
FRONTEND_PID=$!
|
||||
|
||||
echo "==> Starting ngrok tunnel -> $NGROK_URL"
|
||||
"$CONDA_PREFIX/bin/ngrok" http --url="$NGROK_URL" 5299
|
||||
@@ -1,105 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# setup-dreamverse-env.sh — create and configure the dreamverse conda env
|
||||
# from scratch on this aarch64 NFS Slurm cluster.
|
||||
#
|
||||
# Run from the login node (from the repo root):
|
||||
# bash apps/dreamverse/scripts/setup-dreamverse-env.sh
|
||||
#
|
||||
# After this script completes, use launch-dreamverse.sh on a compute node.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(git rev-parse --show-toplevel)"
|
||||
ENV_NAME="dreamverse"
|
||||
LOCAL_DIR="/mnt/local/hal-kevin" # cache/pkgs — keep on local disk
|
||||
CONDA_PREFIX="$HOME/miniconda3/envs/$ENV_NAME" # env — on shared NFS so it survives node changes
|
||||
|
||||
echo "==> Removing existing env if present"
|
||||
conda env remove -p "$CONDA_PREFIX" -y 2>/dev/null || true
|
||||
rm -rf "$CONDA_PREFIX" 2>/dev/null || true
|
||||
|
||||
echo "==> Creating conda env at $CONDA_PREFIX"
|
||||
CONDA_PKGS_DIRS="$LOCAL_DIR/conda/pkgs" conda create -p "$CONDA_PREFIX" python=3.11 -y
|
||||
|
||||
GXX="$CONDA_PREFIX/bin/aarch64-conda-linux-gnu-g++"
|
||||
GCC="$CONDA_PREFIX/bin/aarch64-conda-linux-gnu-gcc"
|
||||
|
||||
echo "==> Installing compiler"
|
||||
CONDA_PKGS_DIRS="$LOCAL_DIR/conda/pkgs" conda install -p "$CONDA_PREFIX" gxx_linux-aarch64 -y
|
||||
|
||||
echo "==> Installing CUDA toolkit (nvcc + headers)"
|
||||
CONDA_PKGS_DIRS="$LOCAL_DIR/conda/pkgs" conda install -p "$CONDA_PREFIX" -c nvidia cuda-toolkit -y
|
||||
|
||||
echo "==> Hiding conflicting libcudart.so.13 immediately"
|
||||
mkdir -p "$CONDA_PREFIX/lib/hidden"
|
||||
mv "$CONDA_PREFIX"/lib/libcudart.so* "$CONDA_PREFIX/lib/hidden/" 2>/dev/null || true
|
||||
|
||||
echo "==> Fixing compiler_compat symlinks"
|
||||
mkdir -p "$CONDA_PREFIX/compiler_compat"
|
||||
ln -sf "$GXX" "$CONDA_PREFIX/compiler_compat/g++"
|
||||
ln -sf "$GCC" "$CONDA_PREFIX/compiler_compat/gcc"
|
||||
|
||||
echo "==> Symlinking CUDA headers to standard location"
|
||||
for f in "$CONDA_PREFIX/targets/sbsa-linux/include/"*; do
|
||||
ln -sf "$f" "$CONDA_PREFIX/include/$(basename "$f")" 2>/dev/null || true
|
||||
done
|
||||
|
||||
echo "==> Installing ffmpeg (native build with x264 + NVENC)"
|
||||
CUDA_PREFIX="$CONDA_PREFIX" bash "$REPO_ROOT/apps/dreamverse/scripts/install_native_ffmpeg.sh"
|
||||
|
||||
echo "==> Installing pip and uv"
|
||||
CONDA_PKGS_DIRS="$LOCAL_DIR/conda/pkgs" conda install -p "$CONDA_PREFIX" pip -y
|
||||
"$CONDA_PREFIX/bin/pip" install uv
|
||||
|
||||
echo "==> Setting compiler env vars"
|
||||
export UV_CACHE_DIR="$LOCAL_DIR/cache"
|
||||
export UV_LINK_MODE=copy
|
||||
export CXX="$CONDA_PREFIX/compiler_compat/g++"
|
||||
export CC="$CONDA_PREFIX/compiler_compat/gcc"
|
||||
export CUDAHOSTCXX="$GXX"
|
||||
export NVCC_PREPEND_FLAGS="-ccbin $GXX -allow-unsupported-compiler"
|
||||
export CUDA_HOME="$CONDA_PREFIX"
|
||||
# Only build for GB200 (sm_100a); CUDA 13 dropped support for older archs
|
||||
export TORCH_CUDA_ARCH_LIST="10.0a"
|
||||
|
||||
echo "==> Installing torch with CUDA 12.8"
|
||||
UV_CACHE_DIR="$LOCAL_DIR/cache" "$CONDA_PREFIX/bin/uv" pip install torch==2.11.0 torchvision \
|
||||
--index-url https://download.pytorch.org/whl/cu128
|
||||
|
||||
echo "==> Hiding any newly introduced libcudart.so.13"
|
||||
mv "$CONDA_PREFIX"/lib/libcudart.so* "$CONDA_PREFIX/lib/hidden/" 2>/dev/null || true
|
||||
|
||||
# Set paths now that torch (and its nvidia packages) are installed
|
||||
CUDA_RT_DIR="$CONDA_PREFIX/lib/python3.11/site-packages/nvidia/cuda_runtime/lib"
|
||||
CUDA_RT_SO="$(ls "$CUDA_RT_DIR"/libcudart.so.* 2>/dev/null | head -1)"
|
||||
|
||||
# The pip nvidia package only has libcudart.so.12 (versioned), not libcudart.so.
|
||||
# The linker needs the unversioned name to satisfy -lcudart. Create a compat dir.
|
||||
mkdir -p "$CONDA_PREFIX/lib/cuda-compat"
|
||||
ln -sf "$CUDA_RT_SO" "$CONDA_PREFIX/lib/cuda-compat/libcudart.so"
|
||||
|
||||
export LIBRARY_PATH="$CONDA_PREFIX/lib/cuda-compat:$CONDA_PREFIX/lib/stubs"
|
||||
export CMAKE_ARGS="-DCUDA_CUDART_LIBRARY=$CUDA_RT_SO -DCUDA_INCLUDE_DIRS=$CONDA_PREFIX/targets/sbsa-linux/include"
|
||||
|
||||
echo "==> Installing build tools"
|
||||
"$CONDA_PREFIX/bin/pip" install scikit-build-core cmake ninja
|
||||
|
||||
echo "==> Initializing git submodules"
|
||||
cd "$REPO_ROOT"
|
||||
git submodule update --init fastvideo-kernel/include/cutlass fastvideo-kernel/include/tk
|
||||
|
||||
echo "==> Building fastvideo-kernel from local source"
|
||||
UV_CACHE_DIR="$LOCAL_DIR/cache" "$CONDA_PREFIX/bin/uv" pip install \
|
||||
-e "./fastvideo-kernel" --no-build-isolation
|
||||
|
||||
echo "==> Installing fastvideo + dreamverse extras"
|
||||
UV_CACHE_DIR="$LOCAL_DIR/cache" "$CONDA_PREFIX/bin/uv" pip install \
|
||||
-e ".[dreamverse]" --no-build-isolation
|
||||
|
||||
echo "==> Installing flashinfer-python (pinned, must be last)"
|
||||
UV_CACHE_DIR="$LOCAL_DIR/cache" "$CONDA_PREFIX/bin/uv" pip install \
|
||||
https://github.com/flashinfer-ai/flashinfer/releases/download/v0.6.11.post3/flashinfer_python-0.6.11.post3-py3-none-any.whl
|
||||
|
||||
echo ""
|
||||
echo "Done. On a compute node run (GPUs 0-3 by default; set CUDA_VISIBLE_DEVICES to restrict):"
|
||||
echo " bash apps/dreamverse/scripts/launch-dreamverse.sh"
|
||||
@@ -110,7 +110,7 @@ default_request:
|
||||
fps: 24 # internal: gpu_pool.py:85 TARGET_FPS
|
||||
|
||||
streaming:
|
||||
# internal: config.py SESSION_TIMEOUT_SECONDS (env DREAMVERSE_SESSION_TIMEOUT_SECONDS, default 300)
|
||||
# internal: config.py:33 SESSION_TIMEOUT_SECONDS = 300
|
||||
session_timeout_seconds: 300
|
||||
# internal: config.py:282-284 GENERATION_SEGMENT_CAP default 6
|
||||
generation_segment_cap: 6
|
||||
|
||||
@@ -9,7 +9,7 @@ import SessionTimeoutModal from "@/components/SessionTimeoutModal";
|
||||
import Sidebar from "@/components/Sidebar";
|
||||
import Header from "@/components/Header";
|
||||
import VideoPlayer from "@/components/VideoPlayer";
|
||||
import Workspace, { SceneHistoryList } from "@/components/Workspace";
|
||||
import Workspace from "@/components/Workspace";
|
||||
import { saveProject, saveProjectMetadata, listProjects, loadProjectClips, deleteProject, pruneOldProjects, type StoredProject, type StoredClip } from "@/lib/projectStorage";
|
||||
import { isInfrastructureError } from "@/lib/ws/reducer";
|
||||
import { useStore } from "@/hooks/useStore";
|
||||
@@ -31,7 +31,7 @@ import { applyNormalizedSocketEvent } from "@/lib/ws/reducer";
|
||||
import { createPromptWindowStore } from "@/stores/promptWindow";
|
||||
import { createRewriteStore } from "@/stores/rewrite";
|
||||
import { createSessionStore } from "@/stores/session";
|
||||
import { createStreamStore, USER_PROMPT_SOURCES } from "@/stores/stream";
|
||||
import { createStreamStore } from "@/stores/stream";
|
||||
import { createUiStore } from "@/stores/ui";
|
||||
import { Button } from "@/components/ui/button";
|
||||
|
||||
@@ -80,7 +80,7 @@ function yieldToEventLoop(): Promise<void> {
|
||||
|
||||
const HERO_WAVE_LIGHT = ["#2A4A98", "#4878E5", "#6FA0F2", "#B0BCC8", "#E8D99E", "#D8C844", "#C2A620"];
|
||||
const HERO_WAVE_DARK = ["#143468", "#1E58B8", "#3892F0", "#80B8E8", "#B8D0EA", "#E2D498", "#DABB50"];
|
||||
const HERO_TEXT = "Direct scenes in seconds with";
|
||||
const HERO_TEXT = "Direct scenes in seconds";
|
||||
|
||||
function HeroTagline() {
|
||||
const ref = useRef<HTMLHeadingElement>(null);
|
||||
@@ -166,8 +166,6 @@ function HeroTagline() {
|
||||
</span>
|
||||
</Fragment>
|
||||
))}
|
||||
<span data-char className="transition-[color,filter] duration-150">{" "}</span>
|
||||
<img src="/logo.svg" alt="FastVideo" className="inline-block h-[1.1em] w-auto align-middle" />
|
||||
</h1>
|
||||
);
|
||||
}
|
||||
@@ -239,7 +237,6 @@ export default function Page() {
|
||||
enhancementEnabled,
|
||||
promptExtensionError,
|
||||
autoExtensionEnabled,
|
||||
manualContinuationMode,
|
||||
autoExtensionTimeoutHint,
|
||||
loopGenerationEnabled,
|
||||
generationPaused,
|
||||
@@ -252,8 +249,6 @@ export default function Page() {
|
||||
livePromptRewriteMode,
|
||||
sessionExpired,
|
||||
projectResetPending,
|
||||
waitingForSegmentPrompt,
|
||||
generatingNextScene,
|
||||
} = sessionState;
|
||||
|
||||
const {
|
||||
@@ -328,11 +323,6 @@ export default function Page() {
|
||||
const [ttffValueMs, setTtffValueMs] = useState<number | null>(null);
|
||||
const ttffIntervalRef = useRef<ReturnType<typeof setInterval> | null>(null);
|
||||
const pendingInitialPromptRef = useRef("");
|
||||
// Prompt id the opening scene is recorded under; sent as initial_rollout_prompt_id
|
||||
// so the backend's pre-seeded opening PromptSubmission emits status updates
|
||||
// (prompt_enhancing/prompt_ready/prompt_fallback_used) against the same id.
|
||||
const pendingInitialPromptIdRef = useRef("");
|
||||
const [initialImageDataUrl, setInitialImageDataUrl] = useState("");
|
||||
const lastArchivedReplayKeyRef = useRef("");
|
||||
const [sidebarOpen, setSidebarOpen] = useState(false);
|
||||
const [currentThumbnail, setCurrentThumbnail] = useState<string | null>(null);
|
||||
@@ -432,11 +422,6 @@ export default function Page() {
|
||||
if (String(e?.source || "") === "user_rewrite" && typeof e?.text === "string" && e.text.trim()) {
|
||||
return e.text.trim();
|
||||
}
|
||||
// Steering opening: the backend overwrites text/source with the enhanced
|
||||
// prompt once ready, so fall back to the stable rawText record.
|
||||
if (e?.steeringUserPrompt && typeof e?.rawText === "string" && e.rawText.trim()) {
|
||||
return e.rawText.trim();
|
||||
}
|
||||
}
|
||||
return "Untitled project";
|
||||
}, [selectedPreset, promptEvents]);
|
||||
@@ -447,39 +432,11 @@ export default function Page() {
|
||||
const canDownloadVideo = useMemo(() => {
|
||||
const currentActiveClip = activeClip as Record<string, any> | null;
|
||||
if (currentActiveClip?.blob instanceof Blob) return true;
|
||||
if ((completedClips as Record<string, any>[]).some((clip) => clip?.blob instanceof Blob)) return true;
|
||||
// Steering only: once playback has started the live AV pipeline holds playable segments, so
|
||||
// the user can download the in-progress video at any time (handleDownloadVideo remuxes live
|
||||
// segments). Auto mode keeps its original blob-gated behavior.
|
||||
return Boolean(manualContinuationMode) && Boolean(avPlaybackStarted);
|
||||
}, [activeClip, completedClips, avPlaybackStarted, manualContinuationMode]);
|
||||
|
||||
// Steering mode scene list (oldest first). Primary source is the user's own words, captured
|
||||
// stably at submit time as `rawText` (the backend later overwrites text/source with the
|
||||
// enhanced prompt, so we never read those). A segment with no user prompt — e.g. a preset's
|
||||
// opening scene — falls back to promptHistory (the actual prompt that drove that segment).
|
||||
const steeringScenes = useMemo(() => {
|
||||
if (!manualContinuationMode) return [] as Record<string, any>[];
|
||||
const userScenes = (promptEvents as Record<string, any>[])
|
||||
.filter((e) => e?.steeringUserPrompt && !e?.steeringFailed && typeof e?.rawText === "string" && e.rawText.trim())
|
||||
.slice()
|
||||
.reverse() // oldest -> newest
|
||||
.map((e) => ({ id: e.promptId, prompt: e.rawText as string }));
|
||||
const scenes: Record<string, any>[] = [];
|
||||
// Preset opening segments: curated seeds with no user prompt of their own.
|
||||
const curatedHists = (promptHistory as Record<string, any>[])
|
||||
.slice()
|
||||
.reverse() // oldest first
|
||||
.filter((h) => !USER_PROMPT_SOURCES.has(String(h?.source || "")) && typeof h?.prompt === "string" && (h.prompt as string).trim());
|
||||
scenes.push(...curatedHists.map((h) => ({ id: h.id || "scene_open", prompt: h.prompt })));
|
||||
scenes.push(...userScenes);
|
||||
return scenes;
|
||||
}, [manualContinuationMode, promptEvents, promptHistory]);
|
||||
return (completedClips as Record<string, any>[]).some((clip) => clip?.blob instanceof Blob);
|
||||
}, [activeClip, completedClips]);
|
||||
|
||||
const hasEdits = useMemo(
|
||||
() => Boolean(sessionStarted) && (
|
||||
(promptEvents as Record<string, any>[]).some((e) => typeof e?.text === "string" && e.text.trim() && String(e?.source || "").trim() === "user_rewrite")
|
||||
),
|
||||
() => Boolean(sessionStarted) && (promptEvents as Record<string, any>[]).some((e) => typeof e?.text === "string" && e.text.trim() && String(e?.source || "").trim() === "user_rewrite"),
|
||||
[sessionStarted, promptEvents],
|
||||
);
|
||||
|
||||
@@ -1464,8 +1421,7 @@ export default function Page() {
|
||||
if (!prompt) return;
|
||||
lastSubmitTimeRef.current = now;
|
||||
const rCAR = !uiStore.get().devtoolsMode && !uiStore.get().demoMode;
|
||||
const inManualMode = sessionStore.get().manualContinuationMode;
|
||||
if (!inManualMode && (rCAR || (sessionStore.get().livePromptRewriteMode && !uiStore.get().demoMode))) {
|
||||
if (rCAR || (sessionStore.get().livePromptRewriteMode && !uiStore.get().demoMode)) {
|
||||
if (rewriteStore.get().rewritingSeedPrompts) return;
|
||||
const rewriteSourcePromptWindowPrompts = getActivePromptWindowPrompts();
|
||||
const nextPendingClip = {
|
||||
@@ -1506,11 +1462,6 @@ export default function Page() {
|
||||
status: "submitted",
|
||||
source: "user_raw",
|
||||
text: prompt,
|
||||
// Stable record of the user's own words for the steering scene list. The backend
|
||||
// later overwrites `text`/`source` with the enhanced prompt via prompt/ready, but
|
||||
// these two fields are never touched by trackPromptEvent.
|
||||
steeringUserPrompt: true,
|
||||
rawText: prompt,
|
||||
});
|
||||
ws.send(
|
||||
JSON.stringify({
|
||||
@@ -1524,14 +1475,7 @@ export default function Page() {
|
||||
activeClipId: shouldUseArchivedPlaybackFallback() ? streamStore.get().activeClipId : "",
|
||||
activePlaybackStartTime: shouldUseArchivedPlaybackFallback() ? streamStore.get().activePlaybackStartTime : 0,
|
||||
});
|
||||
sessionStore.patch({
|
||||
livePromptDraft: "",
|
||||
waitingForSegmentPrompt: false,
|
||||
sessionNotice: "",
|
||||
// Light the "Generating next scene" overlay immediately on a real submit;
|
||||
// stream/media_init (or a fallback/error) clears it.
|
||||
...(inManualMode ? { generatingNextScene: true } : {}),
|
||||
});
|
||||
sessionStore.patch({ livePromptDraft: "" });
|
||||
}
|
||||
|
||||
function setLivePromptRewriteMode(enabled: boolean) {
|
||||
@@ -1606,14 +1550,6 @@ export default function Page() {
|
||||
);
|
||||
}
|
||||
|
||||
// Steering (manual continuation) vs the automatic 6-segment rollout — a pre-session
|
||||
// preference honored when the session starts.
|
||||
function handleManualContinuationToggle(event: any) {
|
||||
sessionStore.patch({
|
||||
manualContinuationMode: Boolean(event.currentTarget.checked),
|
||||
});
|
||||
}
|
||||
|
||||
function handleLoopGenerationToggle(event: any) {
|
||||
sessionStore.patch({
|
||||
loopGenerationEnabled: Boolean(event.currentTarget.checked),
|
||||
@@ -1789,9 +1725,6 @@ export default function Page() {
|
||||
sessionNotice: preserveSessionNotice ? sessionStore.get().sessionNotice : "",
|
||||
sessionExpired: preserveSessionNotice ? sessionStore.get().sessionExpired : false,
|
||||
projectResetPending: false,
|
||||
manualContinuationMode: true,
|
||||
waitingForSegmentPrompt: false,
|
||||
generatingNextScene: false,
|
||||
});
|
||||
rewriteStore.resetSessionState();
|
||||
streamStore.resetSessionState();
|
||||
@@ -1804,7 +1737,6 @@ export default function Page() {
|
||||
function resetToProjectLobbyState() {
|
||||
setVideoMuted(true);
|
||||
pendingInitialPromptRef.current = "";
|
||||
pendingInitialPromptIdRef.current = "";
|
||||
sessionStore.patch({
|
||||
sessionStarted: false,
|
||||
livePromptDraft: "",
|
||||
@@ -1818,9 +1750,6 @@ export default function Page() {
|
||||
sessionNotice: "",
|
||||
sessionExpired: false,
|
||||
projectResetPending: false,
|
||||
manualContinuationMode: true,
|
||||
waitingForSegmentPrompt: false,
|
||||
generatingNextScene: false,
|
||||
});
|
||||
rewriteStore.resetSessionState();
|
||||
streamStore.resetSessionState();
|
||||
@@ -1831,14 +1760,7 @@ export default function Page() {
|
||||
}
|
||||
|
||||
function buildProjectInitPayload(type: "session_init_v2" | "project_init_v1") {
|
||||
const manualMode = Boolean(sessionStore.get().manualContinuationMode);
|
||||
let segmentPrompts = getSessionInitPrompts();
|
||||
// Steering mode: seed the first 2 segments from the preset so there's no
|
||||
// gap between segment 1 and 2; the user drives every subsequent segment.
|
||||
// Force auto/loop off so the backend waits after the seeded prompts run out.
|
||||
if (manualMode) {
|
||||
segmentPrompts = segmentPrompts.slice(0, 2);
|
||||
}
|
||||
const segmentPrompts = getSessionInitPrompts();
|
||||
setSeedPrompts(segmentPrompts);
|
||||
return {
|
||||
type,
|
||||
@@ -1846,17 +1768,11 @@ export default function Page() {
|
||||
preset_label: getInitialPresetLabel(),
|
||||
curated_prompts: segmentPrompts,
|
||||
initial_rollout_prompt: normalizeInitialPrompt(pendingInitialPromptRef.current),
|
||||
// Ties the backend's pre-seeded opening PromptSubmission to the prompt event
|
||||
// recorded in beginProjectLocally so its status updates land on that record.
|
||||
initial_rollout_prompt_id: pendingInitialPromptIdRef.current,
|
||||
initial_image: initialImageDataUrl
|
||||
? { data_url: initialImageDataUrl, mime_type: initialImageDataUrl.split(";")[0].split(":")[1] || "image/png", name: "upload.png" }
|
||||
: null,
|
||||
initial_image: null,
|
||||
single_clip_mode: false,
|
||||
enhancement_enabled: sessionStore.get().enhancementEnabled,
|
||||
auto_extension_enabled: manualMode ? false : sessionStore.get().autoExtensionEnabled,
|
||||
loop_generation_enabled: manualMode ? false : sessionStore.get().loopGenerationEnabled,
|
||||
manual_continuation_mode: manualMode,
|
||||
auto_extension_enabled: sessionStore.get().autoExtensionEnabled,
|
||||
loop_generation_enabled: sessionStore.get().loopGenerationEnabled,
|
||||
};
|
||||
}
|
||||
|
||||
@@ -2035,11 +1951,7 @@ export default function Page() {
|
||||
setCurrentThumbnail(null);
|
||||
const initialPrompt = normalizeInitialPrompt(sessionStore.get().livePromptDraft as string);
|
||||
pendingInitialPromptRef.current = initialPrompt;
|
||||
setInitialImageDataUrl("");
|
||||
const rCAR = !uiStore.get().devtoolsMode && !uiStore.get().demoMode;
|
||||
// The "Steering mode" toggle is authoritative: checked = manual per-segment steering,
|
||||
// unchecked = automatic 6-segment rollout (default).
|
||||
const nextManualContinuationMode = Boolean(sessionStore.get().manualContinuationMode);
|
||||
sessionStore.patch({
|
||||
sessionNotice: "",
|
||||
sessionExpired: false,
|
||||
@@ -2054,9 +1966,6 @@ export default function Page() {
|
||||
autoExtensionTimeoutHint: "",
|
||||
generationPaused: false,
|
||||
projectResetPending: false,
|
||||
manualContinuationMode: nextManualContinuationMode,
|
||||
waitingForSegmentPrompt: false,
|
||||
generatingNextScene: false,
|
||||
});
|
||||
resetPlaybackState();
|
||||
streamStore.patch({
|
||||
@@ -2070,15 +1979,12 @@ export default function Page() {
|
||||
selectedHistoryId: "",
|
||||
});
|
||||
rewriteStore.resetSessionState();
|
||||
pendingInitialPromptIdRef.current = initialPrompt ? makePromptId() : "";
|
||||
if (initialPrompt) {
|
||||
addPromptEvent({
|
||||
promptId: pendingInitialPromptIdRef.current,
|
||||
promptId: makePromptId(),
|
||||
status: "rewrite_requested",
|
||||
source: "user_rewrite",
|
||||
text: initialPrompt,
|
||||
// In steering mode the typed opening is the user's Scene 1 — record it stably.
|
||||
...(nextManualContinuationMode ? { steeringUserPrompt: true, rawText: initialPrompt } : {}),
|
||||
});
|
||||
}
|
||||
setSeedPrompts(getSessionInitPrompts());
|
||||
@@ -2594,7 +2500,6 @@ export default function Page() {
|
||||
selectedPresetId={selectedPresetId as string}
|
||||
enhancementEnabled={enhancementEnabled as boolean}
|
||||
autoExtensionEnabled={autoExtensionEnabled as boolean}
|
||||
manualContinuationEnabled={manualContinuationMode as boolean}
|
||||
loopGenerationEnabled={loopGenerationEnabled as boolean}
|
||||
canJoinSession={canJoinSession as boolean}
|
||||
canSubmitContinuation={canSubmitContinuation}
|
||||
@@ -2607,7 +2512,6 @@ export default function Page() {
|
||||
onEnhancementToggle={handleEnhancementToggle}
|
||||
onCuratedPromptLimitChange={handleCuratedPromptLimitChange}
|
||||
onAutoExtensionToggle={handleAutoExtensionToggle}
|
||||
onManualContinuationToggle={handleManualContinuationToggle}
|
||||
onLoopToggle={handleLoopGenerationToggle}
|
||||
onJoin={joinSession}
|
||||
onLeave={leaveSession}
|
||||
@@ -2737,7 +2641,6 @@ export default function Page() {
|
||||
/>
|
||||
<Header timeLeft={headerTimeLeft} formatTime={formatTime} onToggleSidebar={() => setSidebarOpen((prev) => !prev)} />
|
||||
|
||||
<div className={cn("flex flex-1 min-h-0 flex-col", sessionStarted && "pb-16 sm:pb-28")}>
|
||||
<div className="relative flex flex-1 min-h-0 flex-col justify-center px-4 pb-2 sm:px-6 sm:pb-12">
|
||||
{isViewingMode && (
|
||||
<>
|
||||
@@ -2821,8 +2724,6 @@ export default function Page() {
|
||||
showLivePlayback={showLivePlayback}
|
||||
defaultMuted={videoMuted}
|
||||
canDownload={canDownloadVideo}
|
||||
waitingForSegmentPrompt={waitingForSegmentPrompt as boolean}
|
||||
generatingNextScene={generatingNextScene as boolean}
|
||||
onPlaying={markFirstFrameRendered}
|
||||
onDownload={handleDownloadVideo}
|
||||
/>
|
||||
@@ -2833,7 +2734,6 @@ export default function Page() {
|
||||
<section className={cn("mx-auto w-full max-w-2xl", hasEdits && "flex-1 min-h-0 overflow-y-auto")}>
|
||||
<Workspace
|
||||
promptEvents={promptEvents as any[]}
|
||||
manualMode={manualContinuationMode as boolean}
|
||||
currentThumbnail={currentThumbnail}
|
||||
originalLabel={pendingInitialPromptRef.current || (selectedPreset as Record<string, any>)?.label || ""}
|
||||
sessionStarted={sessionStarted as boolean}
|
||||
@@ -2880,17 +2780,11 @@ export default function Page() {
|
||||
isGenerating={loadingAnimation as boolean}
|
||||
storyPresets={storyPresets as any[]}
|
||||
continuationDraft={livePromptDraft as string}
|
||||
manualContinuationEnabled={manualContinuationMode as boolean}
|
||||
onModeChange={(manual) => sessionStore.patch({ manualContinuationMode: manual })}
|
||||
initialImageDataUrl={initialImageDataUrl}
|
||||
onImageUpload={(dataUrl) => setInitialImageDataUrl(dataUrl)}
|
||||
onImageClear={() => setInitialImageDataUrl("")}
|
||||
canJoinSession={canStartSession}
|
||||
canSubmitContinuation={canSubmitContinuation}
|
||||
sessionExpired={sessionExpired as boolean}
|
||||
sessionNotice={sessionNotice as string}
|
||||
projectResetPending={projectResetPending as boolean}
|
||||
waitingForSegmentPrompt={waitingForSegmentPrompt as boolean}
|
||||
onPresetGenerate={handlePresetGenerate}
|
||||
onContinuationInput={handleLivePromptInput}
|
||||
onContinuationKeydown={handleLivePromptKeydown}
|
||||
@@ -2904,12 +2798,6 @@ export default function Page() {
|
||||
</motion.div>
|
||||
</div>
|
||||
</div>
|
||||
{manualContinuationMode && (
|
||||
<div className="px-4 sm:px-6">
|
||||
<SceneHistoryList sceneHistory={steeringScenes as any[]} />
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</main>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -2,16 +2,13 @@
|
||||
|
||||
import React, { useRef, useState, useCallback, useEffect } from "react";
|
||||
import Image from "next/image";
|
||||
import { Film, ArrowUp, X, Loader2, ArrowLeft, ImagePlus } from "lucide-react";
|
||||
import { Film, ArrowUp, X, Loader2, ArrowLeft } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import LeaveSessionModal, { shouldShowLeaveWarning } from "@/components/LeaveSessionModal";
|
||||
import SpeechToTextButton from "@/components/SpeechToTextButton";
|
||||
import { cn } from "@/lib/utils";
|
||||
|
||||
const PROMPT_MAX_LENGTH = 500;
|
||||
// Must match backend session_init_image.py: MAX_SESSION_INIT_IMAGE_BYTES / SUPPORTED_SESSION_INIT_IMAGE_MIME_TYPES.
|
||||
const IMAGE_MAX_BYTES = 15 * 1024 * 1024;
|
||||
const IMAGE_ALLOWED_TYPES = ["image/png", "image/jpeg", "image/webp"];
|
||||
|
||||
interface Props {
|
||||
sessionStarted?: boolean;
|
||||
@@ -25,12 +22,6 @@ interface Props {
|
||||
sessionNotice?: string;
|
||||
projectResetPending?: boolean;
|
||||
viewingReadOnly?: boolean;
|
||||
waitingForSegmentPrompt?: boolean;
|
||||
manualContinuationEnabled?: boolean;
|
||||
onModeChange?: (manual: boolean) => void;
|
||||
initialImageDataUrl?: string;
|
||||
onImageUpload?: (dataUrl: string, mimeType: string, name: string) => void;
|
||||
onImageClear?: () => void;
|
||||
onPresetGenerate?: (presetId: string) => void;
|
||||
onContinuationInput?: (e: React.ChangeEvent<HTMLTextAreaElement>) => void;
|
||||
onContinuationKeydown?: (e: React.KeyboardEvent<HTMLTextAreaElement>) => void;
|
||||
@@ -55,12 +46,6 @@ export default function ChatBar({
|
||||
sessionNotice = "",
|
||||
projectResetPending = false,
|
||||
viewingReadOnly = false,
|
||||
waitingForSegmentPrompt = false,
|
||||
manualContinuationEnabled = false,
|
||||
onModeChange = () => {},
|
||||
initialImageDataUrl = "",
|
||||
onImageUpload = () => {},
|
||||
onImageClear = () => {},
|
||||
onPresetGenerate = () => {},
|
||||
onContinuationInput = () => {},
|
||||
onContinuationKeydown = () => {},
|
||||
@@ -74,51 +59,15 @@ export default function ChatBar({
|
||||
}: Props) {
|
||||
const [sttBusy, setSttBusy] = useState(false);
|
||||
const [leaveModalOpen, setLeaveModalOpen] = useState(false);
|
||||
const [imageError, setImageError] = useState("");
|
||||
const fileInputRef = useRef<HTMLInputElement>(null);
|
||||
|
||||
const processImageFile = useCallback((file: File) => {
|
||||
if (!IMAGE_ALLOWED_TYPES.includes(file.type)) {
|
||||
setImageError("Unsupported image type. Use a PNG, JPEG, or WebP image.");
|
||||
return;
|
||||
}
|
||||
if (file.size > IMAGE_MAX_BYTES) {
|
||||
setImageError("Image is too large. The maximum size is 15MB.");
|
||||
return;
|
||||
}
|
||||
const reader = new FileReader();
|
||||
reader.onload = (e) => {
|
||||
const dataUrl = e.target?.result as string;
|
||||
if (dataUrl) {
|
||||
setImageError("");
|
||||
onImageUpload(dataUrl, file.type, file.name);
|
||||
}
|
||||
};
|
||||
reader.onerror = () => {
|
||||
setImageError("Could not read the image file. Please try again.");
|
||||
};
|
||||
reader.readAsDataURL(file);
|
||||
}, [onImageUpload]);
|
||||
|
||||
const handleImagePaste = useCallback((e: React.ClipboardEvent) => {
|
||||
if (sessionStarted) return;
|
||||
const items = Array.from(e.clipboardData?.items ?? []);
|
||||
const imageItem = items.find((item) => item.type.startsWith("image/"));
|
||||
if (!imageItem) return;
|
||||
const file = imageItem.getAsFile();
|
||||
if (file) processImageFile(file);
|
||||
}, [sessionStarted, processImageFile]);
|
||||
const showSpinner = isGenerating || rewritingSeedPrompts;
|
||||
const isBusy = isGenerating || rewritingSeedPrompts || projectResetPending;
|
||||
const messagePlaceholder = projectResetPending
|
||||
? "Starting new project\u2026"
|
||||
: isBusy
|
||||
? "Generating video\u2026"
|
||||
: waitingForSegmentPrompt
|
||||
? "Describe the next scene\u2026"
|
||||
: !sessionStarted
|
||||
? "What video are you imagining?"
|
||||
: "What do you want to edit?";
|
||||
: !sessionStarted
|
||||
? "What video are you imagining?"
|
||||
: "What do you want to edit?";
|
||||
const actionLabel = !sessionStarted ? "Generate" : "Rewrite rollout";
|
||||
|
||||
const inputRef = useRef<HTMLTextAreaElement>(null);
|
||||
@@ -318,9 +267,9 @@ export default function ChatBar({
|
||||
<Button onClick={onStartNewProject} size="sm" className="rounded-full px-5">
|
||||
New Project
|
||||
</Button>
|
||||
<a href="https://haoailab.com/blogs/dreamverse/" target="_blank" rel="noopener noreferrer">
|
||||
<a href="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
|
||||
<Button variant="outline" size="sm" className="rounded-full px-5">
|
||||
Blog
|
||||
Join Waitlist
|
||||
</Button>
|
||||
</a>
|
||||
</div>
|
||||
@@ -397,31 +346,6 @@ export default function ChatBar({
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
{!sessionStarted && imageError && (
|
||||
<div className="rounded-xl border border-rose-500/20 bg-rose-500/10 px-4 py-2.5 text-center text-xs text-rose-700 dark:text-rose-300">
|
||||
{imageError}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{!sessionStarted && initialImageDataUrl && (
|
||||
<div className="flex items-center gap-2 rounded-2xl border border-input bg-card/65 px-3 py-2">
|
||||
<img src={initialImageDataUrl} alt="Initial frame" className="h-12 w-12 rounded-lg object-cover" />
|
||||
<span className="flex-1 truncate text-xs text-muted-foreground">Starting image set</span>
|
||||
<button type="button" onClick={() => { setImageError(""); onImageClear(); }} className="text-muted-foreground hover:text-foreground transition-colors">
|
||||
<X className="size-4" />
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept={IMAGE_ALLOWED_TYPES.join(",")}
|
||||
className="hidden"
|
||||
onChange={(e) => { const f = e.target.files?.[0]; if (f) processImageFile(f); e.target.value = ""; }}
|
||||
/>
|
||||
|
||||
<div
|
||||
className={cn(
|
||||
"flex min-w-0 items-center gap-1.5 rounded-4xl border py-2.5 pl-5 pr-2.5 shadow-md backdrop-blur-sm transition-all duration-200",
|
||||
@@ -435,7 +359,6 @@ export default function ChatBar({
|
||||
value={continuationDraft}
|
||||
onChange={onContinuationInput}
|
||||
onKeyDown={handleKeyDown}
|
||||
onPaste={handleImagePaste}
|
||||
placeholder={sttBusy ? "Listening\u2026" : messagePlaceholder}
|
||||
maxLength={PROMPT_MAX_LENGTH}
|
||||
disabled={isBusy || sttBusy}
|
||||
@@ -446,19 +369,6 @@ export default function ChatBar({
|
||||
)}
|
||||
/>
|
||||
{onSpeechTranscript && <SpeechToTextButton disabled={isBusy} onTranscript={onSpeechTranscript} onInterimChange={onSpeechInterimChange} onBusyChange={setSttBusy} />}
|
||||
{!sessionStarted && (
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="icon-sm"
|
||||
title="Add starting image"
|
||||
onClick={() => fileInputRef.current?.click()}
|
||||
disabled={isBusy}
|
||||
className="shrink-0 rounded-full text-muted-foreground hover:text-foreground"
|
||||
>
|
||||
<ImagePlus className="size-4" />
|
||||
</Button>
|
||||
)}
|
||||
{!sessionStarted ? (
|
||||
<Button
|
||||
aria-label={actionLabel}
|
||||
|
||||
@@ -8,6 +8,7 @@ import { Badge } from "@/components/ui/badge";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { ThemeToggle } from "@/components/ui/theme-toggle";
|
||||
|
||||
const FASTVIDEO_REPO_URL = "https://haoailab.com/blogs/dreamverse/";
|
||||
const FASTVIDEO_BLOG_URL = "https://haoailab.com/blogs/dreamverse/";
|
||||
|
||||
interface Props {
|
||||
@@ -28,13 +29,13 @@ export default function Header({ timeLeft = null, formatTime = (seconds) => `${s
|
||||
<SidePanelOpenFilled size={20} />
|
||||
</Button>
|
||||
)}
|
||||
<a href="/" title="FastVideo home">
|
||||
<a href={FASTVIDEO_REPO_URL} target="_blank" rel="noopener noreferrer" title="FastVideo on GitHub">
|
||||
<Image src="/logo.svg" alt="FastVideo" width={32} height={32} className="h-8 w-auto sm:h-9 transition-opacity hover:opacity-70" />
|
||||
</a>
|
||||
<div className="hidden sm:flex items-center gap-3">
|
||||
<a href={FASTVIDEO_BLOG_URL} target="_blank" rel="noopener noreferrer">
|
||||
<a href="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
|
||||
<Button variant="outline" size="sm" className="gap-1.5 rounded-full px-3 text-xs">
|
||||
Blog
|
||||
Join Waitlist
|
||||
<ExternalLink className="size-3 opacity-60" />
|
||||
</Button>
|
||||
</a>
|
||||
@@ -52,9 +53,9 @@ export default function Header({ timeLeft = null, formatTime = (seconds) => `${s
|
||||
</div>
|
||||
|
||||
<div className="flex sm:hidden items-center gap-2 px-4 pb-3">
|
||||
<a href={FASTVIDEO_BLOG_URL} target="_blank" rel="noopener noreferrer">
|
||||
<a href="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
|
||||
<Button variant="outline" size="sm" className="gap-1.5 rounded-full px-3 text-xs">
|
||||
Blog
|
||||
Join Waitlist
|
||||
<ExternalLink className="size-3 opacity-60" />
|
||||
</Button>
|
||||
</a>
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
import React, { useState, useEffect, useRef, useCallback } from "react";
|
||||
import { cn } from "@/lib/utils";
|
||||
import { PlayFilledAlt } from "@carbon/icons-react";
|
||||
import { Check, ChevronDown, Download, Loader2, Share } from "lucide-react";
|
||||
import { Download, Loader2, Share } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
interface VideoPlayerProps {
|
||||
videoRef?: React.RefCallback<HTMLVideoElement>;
|
||||
@@ -21,8 +21,6 @@ interface VideoPlayerProps {
|
||||
showLivePlayback?: boolean;
|
||||
defaultMuted?: boolean;
|
||||
rewritePending?: boolean;
|
||||
waitingForSegmentPrompt?: boolean;
|
||||
generatingNextScene?: boolean;
|
||||
onPlaying?: () => void;
|
||||
onDownload?: () => void;
|
||||
}
|
||||
@@ -59,8 +57,6 @@ export default function VideoPlayer({
|
||||
showLivePlayback = true,
|
||||
defaultMuted = true,
|
||||
rewritePending = false,
|
||||
waitingForSegmentPrompt = false,
|
||||
generatingNextScene = false,
|
||||
onPlaying = () => {},
|
||||
onDownload,
|
||||
}: VideoPlayerProps) {
|
||||
@@ -92,137 +88,6 @@ export default function VideoPlayer({
|
||||
setCanShare(typeof navigator.canShare === "function" && window.matchMedia("(pointer: coarse)").matches);
|
||||
}, []);
|
||||
|
||||
// Steering mode: the backend may signal "waiting for next prompt" while the current
|
||||
// segment is still PLAYING (it generates ahead). Only surface the "Segment complete"
|
||||
// overlay once the playhead actually reaches the end of the buffered segment.
|
||||
const [playbackReachedEnd, setPlaybackReachedEnd] = useState(false);
|
||||
// Steering mode: when the user submits the next scene, the segment is generated
|
||||
// (a few seconds of latency) before frames stream. Show a "Generating next scene…"
|
||||
// indicator across that gap so the frozen frame isn't silent. Driven by the explicit
|
||||
// generatingNextScene state (set on scene submit / prompt selection), never inferred
|
||||
// from waitingForSegmentPrompt edges.
|
||||
const [generatingNext, setGeneratingNext] = useState(false);
|
||||
// End of the buffered timeline captured the moment generation starts; the freshly
|
||||
// generated segment extends the buffer past this, which is how we know it landed.
|
||||
const genBoundaryRef = useRef(0);
|
||||
|
||||
// Steering mode: the backend may signal "waiting for next prompt" while the current
|
||||
// segment is still PLAYING (it generates ahead). Track whether the playhead has reached
|
||||
// the end of the buffered segment so end-overlays only show there. Keep tracking through
|
||||
// the generating phase too, so scrubbing back and replaying to the end re-shows them.
|
||||
useEffect(() => {
|
||||
if (!waitingForSegmentPrompt && !generatingNext) {
|
||||
setPlaybackReachedEnd(false);
|
||||
return;
|
||||
}
|
||||
const el = liveVideoEl.current;
|
||||
if (!el) return;
|
||||
const check = () => {
|
||||
try {
|
||||
const buffered = el.buffered;
|
||||
if (buffered.length === 0) return;
|
||||
const end = buffered.end(buffered.length - 1);
|
||||
// Track proximity both ways: scrubbing back off the end hides the overlay,
|
||||
// playing forward to the end re-shows it.
|
||||
setPlaybackReachedEnd(el.ended || end - el.currentTime <= 0.2);
|
||||
} catch {
|
||||
/* buffered access can throw mid-append */
|
||||
}
|
||||
};
|
||||
check();
|
||||
el.addEventListener("timeupdate", check);
|
||||
el.addEventListener("ended", check);
|
||||
el.addEventListener("waiting", check);
|
||||
el.addEventListener("stalled", check);
|
||||
el.addEventListener("pause", check);
|
||||
el.addEventListener("seeking", check);
|
||||
el.addEventListener("seeked", check);
|
||||
el.addEventListener("playing", check);
|
||||
el.addEventListener("progress", check);
|
||||
return () => {
|
||||
el.removeEventListener("timeupdate", check);
|
||||
el.removeEventListener("ended", check);
|
||||
el.removeEventListener("waiting", check);
|
||||
el.removeEventListener("stalled", check);
|
||||
el.removeEventListener("pause", check);
|
||||
el.removeEventListener("seeking", check);
|
||||
el.removeEventListener("seeked", check);
|
||||
el.removeEventListener("playing", check);
|
||||
el.removeEventListener("progress", check);
|
||||
};
|
||||
}, [waitingForSegmentPrompt, generatingNext]);
|
||||
|
||||
useEffect(() => {
|
||||
if (waitingForSegmentPrompt || !sessionStarted) {
|
||||
// Back to waiting (or session over): nothing is generating.
|
||||
setGeneratingNext(false);
|
||||
return;
|
||||
}
|
||||
if (!generatingNextScene) return;
|
||||
// Snapshot the current end of the buffered timeline; the generated segment will
|
||||
// extend the buffer past this boundary.
|
||||
const el = liveVideoEl.current;
|
||||
let boundary = el?.currentTime ?? 0;
|
||||
try {
|
||||
const b = el?.buffered;
|
||||
if (b && b.length) boundary = Math.max(boundary, b.end(b.length - 1));
|
||||
} catch {
|
||||
/* buffered access can throw mid-append */
|
||||
}
|
||||
genBoundaryRef.current = boundary;
|
||||
setGeneratingNext(true);
|
||||
}, [generatingNextScene, waitingForSegmentPrompt, sessionStarted]);
|
||||
|
||||
useEffect(() => {
|
||||
if (!generatingNext) return;
|
||||
if (!sessionStarted) {
|
||||
setGeneratingNext(false);
|
||||
return;
|
||||
}
|
||||
const el = liveVideoEl.current;
|
||||
if (!el) return;
|
||||
// Clear the instant the freshly generated segment lands: the buffer grows past the
|
||||
// boundary captured at generation start (or the playhead advances into the new
|
||||
// frames). Deliberately NOT a bare "playing" handler — scrubbing back and replaying
|
||||
// the EXISTING segment must keep "Generating" up until the new frames actually arrive.
|
||||
const check = () => {
|
||||
try {
|
||||
const b = el.buffered;
|
||||
const end = b.length ? b.end(b.length - 1) : 0;
|
||||
if (end > genBoundaryRef.current + 0.25 || el.currentTime > genBoundaryRef.current + 0.1) {
|
||||
setGeneratingNext(false);
|
||||
}
|
||||
} catch {
|
||||
/* buffered access can throw mid-append */
|
||||
}
|
||||
};
|
||||
check();
|
||||
el.addEventListener("progress", check);
|
||||
el.addEventListener("timeupdate", check);
|
||||
el.addEventListener("durationchange", check);
|
||||
return () => {
|
||||
el.removeEventListener("progress", check);
|
||||
el.removeEventListener("timeupdate", check);
|
||||
el.removeEventListener("durationchange", check);
|
||||
};
|
||||
}, [generatingNext, sessionStarted]);
|
||||
|
||||
// Drive a ~10.5s progress bar during generation so the wait has a visible ETA.
|
||||
const GEN_DURATION_MS = 10500;
|
||||
const [genProgress, setGenProgress] = useState(0);
|
||||
useEffect(() => {
|
||||
if (!generatingNext) {
|
||||
setGenProgress(0);
|
||||
return;
|
||||
}
|
||||
const start = performance.now();
|
||||
setGenProgress(0);
|
||||
const id = setInterval(() => {
|
||||
setGenProgress(Math.min((performance.now() - start) / GEN_DURATION_MS, 1));
|
||||
}, 50);
|
||||
return () => clearInterval(id);
|
||||
}, [generatingNext]);
|
||||
|
||||
return (
|
||||
<div className="mx-auto w-full max-w-3xl mb-2 sm:mb-6">
|
||||
<div className="rounded-2xl border border-border bg-card/50 p-2 shadow-lg backdrop-blur-md">
|
||||
@@ -239,7 +104,7 @@ export default function VideoPlayer({
|
||||
<PlayFilledAlt className="size-10 text-white/25" />
|
||||
<p className="text-sm text-white/50">Your video will appear here</p>
|
||||
</div>
|
||||
) : !avPlaybackStarted && !mediaAppendError && !inQueue && !waitingForSegmentPrompt && loadingAnimation ? (
|
||||
) : !avPlaybackStarted && !mediaAppendError && !inQueue && loadingAnimation ? (
|
||||
<div className="absolute inset-0 flex flex-col items-center justify-center gap-4 bg-slate-900/60 p-4 backdrop-blur-[2px]">
|
||||
<div className="pointer-events-none absolute inset-0 overflow-hidden">
|
||||
<div className="absolute inset-0 -translate-x-full animate-[shimmer_3s_ease-in-out_infinite] bg-gradient-to-r from-transparent via-white/[0.04] to-transparent" />
|
||||
@@ -249,35 +114,6 @@ export default function VideoPlayer({
|
||||
</div>
|
||||
) : null}
|
||||
|
||||
{/* Steering mode: this segment finished — wait gracefully for the user's next scene
|
||||
instead of spinning. The last frame stays visible behind a soft bottom gradient. */}
|
||||
{sessionStarted && waitingForSegmentPrompt && playbackReachedEnd && !mediaAppendError && !inQueue && (
|
||||
<div className="pointer-events-none absolute inset-0 flex flex-col items-center justify-end gap-2 bg-gradient-to-t from-slate-950/85 via-slate-950/15 to-transparent p-5 pb-6 text-center">
|
||||
<div className="flex size-9 items-center justify-center rounded-full border border-white/25 bg-white/10 shadow-lg backdrop-blur-md">
|
||||
<Check className="size-4 text-white/90" />
|
||||
</div>
|
||||
<div className="space-y-0.5">
|
||||
<p className="text-sm font-medium text-white/95">Segment complete</p>
|
||||
<p className="text-xs text-white/65">Describe the next scene below to keep going</p>
|
||||
</div>
|
||||
<ChevronDown className="size-4 animate-bounce text-white/45" />
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Steering mode: generating the next segment — show a ~10.5s progress bar so the wait has an ETA.
|
||||
Gated on playbackReachedEnd like "Segment complete": scrubbing back hides it, playing to the end re-shows it. */}
|
||||
{sessionStarted && generatingNext && playbackReachedEnd && !mediaAppendError && !inQueue && (
|
||||
<div className="pointer-events-none absolute inset-0 flex flex-col items-center justify-end gap-3 bg-gradient-to-t from-slate-950/85 via-slate-950/15 to-transparent p-5 pb-7 text-center">
|
||||
<p className="text-sm font-medium text-white/95">Generating next scene…</p>
|
||||
<div className="h-1.5 w-48 overflow-hidden rounded-full bg-white/15 shadow-sm">
|
||||
<div
|
||||
className="h-full rounded-full bg-white/85 transition-[width] duration-100 ease-linear"
|
||||
style={{ width: `${Math.round(genProgress * 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{rewritePending && avPlaybackStarted && (
|
||||
<div className="absolute inset-x-0 bottom-0 z-10 flex items-center justify-center gap-2 bg-gradient-to-t from-black/60 to-transparent px-4 pb-12 pt-8 pointer-events-none">
|
||||
<Loader2 className="size-4 animate-spin text-white/90" />
|
||||
|
||||
@@ -3,90 +3,11 @@ import React, { useRef, useMemo, useEffect, useCallback, useState } from "react"
|
||||
import { motion, useAnimationControls } from "framer-motion";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { cn } from "@/lib/utils";
|
||||
import { Check, Clapperboard, Lightbulb, Pencil } from "lucide-react";
|
||||
import { Check, Lightbulb, Pencil } from "lucide-react";
|
||||
|
||||
export const WORKSPACE_ORIGINAL_SELECTION_KEY = "original";
|
||||
export const WORKSPACE_CURRENT_SELECTION_KEY = "current";
|
||||
|
||||
export function SceneHistoryList({ sceneHistory = [] }: { sceneHistory?: Record<string, any>[] }) {
|
||||
const bottomSentinelRef = useRef<HTMLDivElement>(null);
|
||||
const topSentinelRef = useRef<HTMLDivElement>(null);
|
||||
const [showTopFade, setShowTopFade] = useState(false);
|
||||
|
||||
const scenes = useMemo(
|
||||
() => (sceneHistory || []).filter((s) => normalizeText(s?.prompt)),
|
||||
[sceneHistory],
|
||||
);
|
||||
|
||||
const scrollToBottom = useCallback(() => {
|
||||
setTimeout(() => {
|
||||
bottomSentinelRef.current?.scrollIntoView({ block: "end", behavior: "smooth" });
|
||||
}, 60);
|
||||
}, []);
|
||||
|
||||
useEffect(() => {
|
||||
if (scenes.length > 0) scrollToBottom();
|
||||
}, [scenes.length, scrollToBottom]);
|
||||
|
||||
useEffect(() => {
|
||||
const sentinel = bottomSentinelRef.current;
|
||||
if (!sentinel || typeof ResizeObserver === "undefined") return;
|
||||
let container: HTMLElement | null = sentinel.parentElement;
|
||||
while (container) {
|
||||
const oy = getComputedStyle(container).overflowY;
|
||||
if (oy === "auto" || oy === "scroll") break;
|
||||
container = container.parentElement;
|
||||
}
|
||||
if (!container) return;
|
||||
const ro = new ResizeObserver(() => {
|
||||
const nearBottom = container!.scrollHeight - container!.scrollTop - container!.clientHeight < 96;
|
||||
if (nearBottom) scrollToBottom();
|
||||
});
|
||||
ro.observe(container);
|
||||
return () => ro.disconnect();
|
||||
}, [scenes.length, scrollToBottom]);
|
||||
|
||||
useEffect(() => {
|
||||
const el = topSentinelRef.current;
|
||||
if (!el) return;
|
||||
const observer = new IntersectionObserver(([entry]) => setShowTopFade(!entry.isIntersecting), { threshold: 0.1 });
|
||||
observer.observe(el);
|
||||
return () => observer.disconnect();
|
||||
}, [scenes.length]);
|
||||
|
||||
if (scenes.length === 0) return null;
|
||||
|
||||
return (
|
||||
<section className="relative z-10 flex flex-col mx-auto w-full max-w-2xl max-h-32 overflow-y-auto">
|
||||
<div
|
||||
className={cn(
|
||||
"pointer-events-none sticky top-0 z-20 -mb-12 h-12 bg-linear-to-b from-background to-transparent transition-opacity duration-200",
|
||||
showTopFade ? "opacity-100" : "opacity-0",
|
||||
)}
|
||||
aria-hidden="true"
|
||||
/>
|
||||
<div ref={topSentinelRef} className="h-0 w-0" aria-hidden="true" />
|
||||
<div className="flex flex-col gap-2 pt-12 pb-4">
|
||||
{scenes.map((scene, index) => (
|
||||
<div
|
||||
key={scene.id || index}
|
||||
className="flex items-start gap-3 rounded-xl p-3 transition-colors duration-200 hover:bg-slate-200/50 hover:dark:bg-slate-800/30"
|
||||
>
|
||||
<div className="flex min-w-0 flex-1 flex-col gap-2">
|
||||
<Badge variant="secondary" className="horizontal gap-2 items-center w-fit">
|
||||
<Clapperboard className="size-3 opacity-70" />
|
||||
{`Scene ${index + 1}`}
|
||||
</Badge>
|
||||
<p className="line-clamp-2 text-sm leading-5 text-muted-foreground">{scene.prompt}</p>
|
||||
</div>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
<div ref={bottomSentinelRef} className="h-0 w-0" aria-hidden="true" />
|
||||
</section>
|
||||
);
|
||||
}
|
||||
|
||||
interface WorkspaceProps {
|
||||
promptEvents?: Record<string, any>[];
|
||||
currentThumbnail?: string | null;
|
||||
@@ -98,7 +19,6 @@ interface WorkspaceProps {
|
||||
selectedClipId?: string;
|
||||
selectedEntryKey?: string;
|
||||
originalClipId?: string;
|
||||
manualMode?: boolean;
|
||||
}
|
||||
|
||||
function normalizeText(value: any): string {
|
||||
@@ -298,7 +218,7 @@ function ChromaGradient({ sessionStarted = false }: { sessionStarted?: boolean }
|
||||
);
|
||||
}
|
||||
|
||||
export default function Workspace({ promptEvents = [], currentThumbnail = null, originalLabel = "", sessionStarted = false, onSelectOriginal, onSelectEvent, onSelectCurrent, selectedClipId, selectedEntryKey: selectedEntryKeyProp, originalClipId = "", manualMode = false }: WorkspaceProps) {
|
||||
export default function Workspace({ promptEvents = [], currentThumbnail = null, originalLabel = "", sessionStarted = false, onSelectOriginal, onSelectEvent, onSelectCurrent, selectedClipId, selectedEntryKey: selectedEntryKeyProp, originalClipId = "" }: WorkspaceProps) {
|
||||
const bottomSentinelRef = useRef<HTMLDivElement>(null);
|
||||
const topSentinelRef = useRef<HTMLDivElement>(null);
|
||||
const [showTopFade, setShowTopFade] = useState(false);
|
||||
@@ -339,14 +259,6 @@ export default function Workspace({ promptEvents = [], currentThumbnail = null,
|
||||
return () => observer.disconnect();
|
||||
}, [conversationEvents.length]);
|
||||
|
||||
if (manualMode) {
|
||||
return (
|
||||
<div className="mt-auto flex flex-col">
|
||||
<ChromaGradient sessionStarted={sessionStarted} />
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="mt-auto flex flex-col">
|
||||
<ChromaGradient sessionStarted={sessionStarted} />
|
||||
|
||||
@@ -34,7 +34,6 @@ interface DevtoolsComposerProps {
|
||||
demoMode?: boolean;
|
||||
enhancementEnabled?: boolean;
|
||||
autoExtensionEnabled?: boolean;
|
||||
manualContinuationEnabled?: boolean;
|
||||
loopGenerationEnabled?: boolean;
|
||||
curatedPromptLimit?: number;
|
||||
maxCuratedPromptCount?: number;
|
||||
@@ -50,7 +49,6 @@ interface DevtoolsComposerProps {
|
||||
onEnhancementToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onCuratedPromptLimitChange?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onAutoExtensionToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onManualContinuationToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onLoopToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onLivePromptModeToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onSpeechTranscript?: (text: string) => void;
|
||||
@@ -73,7 +71,6 @@ export default function DevtoolsComposer({
|
||||
demoMode = false,
|
||||
enhancementEnabled = true,
|
||||
autoExtensionEnabled = false,
|
||||
manualContinuationEnabled = false,
|
||||
loopGenerationEnabled = false,
|
||||
curatedPromptLimit = 0,
|
||||
maxCuratedPromptCount = 0,
|
||||
@@ -89,7 +86,6 @@ export default function DevtoolsComposer({
|
||||
onEnhancementToggle = () => {},
|
||||
onCuratedPromptLimitChange = () => {},
|
||||
onAutoExtensionToggle = () => {},
|
||||
onManualContinuationToggle = () => {},
|
||||
onLoopToggle = () => {},
|
||||
onLivePromptModeToggle = () => {},
|
||||
onSpeechTranscript,
|
||||
@@ -332,28 +328,6 @@ export default function DevtoolsComposer({
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex items-start gap-3">
|
||||
<Checkbox
|
||||
id="devtools-steering-mode"
|
||||
checked={manualContinuationEnabled}
|
||||
onCheckedChange={(checked) =>
|
||||
onManualContinuationToggle({
|
||||
target: { checked: Boolean(checked) },
|
||||
currentTarget: { checked: Boolean(checked) },
|
||||
} as React.ChangeEvent<HTMLInputElement>)
|
||||
}
|
||||
/>
|
||||
<div className="space-y-1">
|
||||
<Label htmlFor="devtools-steering-mode">
|
||||
Steering mode
|
||||
</Label>
|
||||
<p className="text-sm text-muted-foreground">
|
||||
Drive each segment manually — type the next scene to
|
||||
continue (vs the automatic 6-segment rollout).
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex items-start gap-3">
|
||||
<Checkbox
|
||||
id="devtools-loop-generation"
|
||||
|
||||
@@ -21,7 +21,6 @@ interface DevtoolsShellProps {
|
||||
selectedPresetId?: string;
|
||||
enhancementEnabled?: boolean;
|
||||
autoExtensionEnabled?: boolean;
|
||||
manualContinuationEnabled?: boolean;
|
||||
loopGenerationEnabled?: boolean;
|
||||
canJoinSession?: boolean;
|
||||
canSubmitContinuation?: boolean;
|
||||
@@ -35,7 +34,6 @@ interface DevtoolsShellProps {
|
||||
onEnhancementToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onCuratedPromptLimitChange?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onAutoExtensionToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onManualContinuationToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onLoopToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onJoin?: () => void;
|
||||
onLeave?: () => void;
|
||||
@@ -130,7 +128,6 @@ export default function DevtoolsShell({
|
||||
selectedPresetId = '',
|
||||
enhancementEnabled = true,
|
||||
autoExtensionEnabled = false,
|
||||
manualContinuationEnabled = false,
|
||||
loopGenerationEnabled = false,
|
||||
canJoinSession = false,
|
||||
canSubmitContinuation = false,
|
||||
@@ -144,7 +141,6 @@ export default function DevtoolsShell({
|
||||
onEnhancementToggle = () => {},
|
||||
onCuratedPromptLimitChange = () => {},
|
||||
onAutoExtensionToggle = () => {},
|
||||
onManualContinuationToggle = () => {},
|
||||
onLoopToggle = () => {},
|
||||
onJoin = () => {},
|
||||
onLeave = () => {},
|
||||
@@ -288,7 +284,6 @@ export default function DevtoolsShell({
|
||||
demoMode={demoMode}
|
||||
enhancementEnabled={enhancementEnabled}
|
||||
autoExtensionEnabled={autoExtensionEnabled}
|
||||
manualContinuationEnabled={manualContinuationEnabled}
|
||||
loopGenerationEnabled={loopGenerationEnabled}
|
||||
curatedPromptLimit={curatedPromptLimit}
|
||||
maxCuratedPromptCount={maxCuratedPromptCount}
|
||||
@@ -304,7 +299,6 @@ export default function DevtoolsShell({
|
||||
onEnhancementToggle={onEnhancementToggle}
|
||||
onCuratedPromptLimitChange={onCuratedPromptLimitChange}
|
||||
onAutoExtensionToggle={onAutoExtensionToggle}
|
||||
onManualContinuationToggle={onManualContinuationToggle}
|
||||
onLoopToggle={onLoopToggle}
|
||||
onLivePromptModeToggle={onLivePromptModeToggle}
|
||||
onSpeechTranscript={onSpeechTranscript}
|
||||
|
||||
@@ -57,32 +57,4 @@ describe('prependPromptEvent', () => {
|
||||
expect(next[0].promptId).toBe('new');
|
||||
expect(next.some((item: any) => item.promptId === 'p-23')).toBe(false);
|
||||
});
|
||||
|
||||
it('never drops steering scene events when capping', () => {
|
||||
// 30 scenes interleaved with 30 other events — well past the cap.
|
||||
let events: Record<string, any>[] = [];
|
||||
for (let i = 0; i < 30; i += 1) {
|
||||
events = prependPromptEvent(events, {
|
||||
promptId: `scene-${i}`,
|
||||
status: 'submitted',
|
||||
steeringUserPrompt: true,
|
||||
rawText: `scene ${i}`,
|
||||
});
|
||||
events = prependPromptEvent(events, {
|
||||
promptId: `other-${i}`,
|
||||
status: 'submitted',
|
||||
});
|
||||
}
|
||||
|
||||
const scenes = events.filter((e) => e.steeringUserPrompt);
|
||||
expect(scenes).toHaveLength(30);
|
||||
// Oldest-first scene order (and therefore numbering) is stable and complete.
|
||||
expect(scenes[scenes.length - 1].promptId).toBe('scene-0');
|
||||
expect(scenes[0].promptId).toBe('scene-29');
|
||||
// Non-scene events are still capped, oldest dropped first.
|
||||
const others = events.filter((e) => !e.steeringUserPrompt);
|
||||
expect(others.length).toBeLessThanOrEqual(24);
|
||||
expect(others.some((e) => e.promptId === 'other-0')).toBe(false);
|
||||
expect(others[0].promptId).toBe('other-29');
|
||||
});
|
||||
});
|
||||
|
||||
@@ -16,19 +16,5 @@ export function prependPromptEvent(
|
||||
events: Record<string, any>[],
|
||||
event: Record<string, any>,
|
||||
): Record<string, any>[] {
|
||||
const next = [event, ...events];
|
||||
if (next.length <= MAX_PROMPT_EVENTS) {
|
||||
return next;
|
||||
}
|
||||
// Steering scene events (steeringUserPrompt) are exempt from the cap: the
|
||||
// scene list is derived from them and must stay complete and stably numbered
|
||||
// for long sessions. Only the oldest non-scene events are dropped.
|
||||
let nonSceneKept = 0;
|
||||
return next.filter((e) => {
|
||||
if (e?.steeringUserPrompt) {
|
||||
return true;
|
||||
}
|
||||
nonSceneKept += 1;
|
||||
return nonSceneKept <= MAX_PROMPT_EVENTS;
|
||||
});
|
||||
return [event, ...events].slice(0, MAX_PROMPT_EVENTS);
|
||||
}
|
||||
|
||||
@@ -1,11 +1,6 @@
|
||||
import { describe, expect, it } from 'vitest';
|
||||
|
||||
import { applyNormalizedSocketEvent, resolveSessionErrorMessage } from './reducer';
|
||||
import { createSessionStore } from '../../stores/session';
|
||||
import { createRewriteStore } from '../../stores/rewrite';
|
||||
import { createStreamStore } from '../../stores/stream';
|
||||
import { createUiStore } from '../../stores/ui';
|
||||
import { createPromptWindowStore } from '../../stores/promptWindow';
|
||||
import { resolveSessionErrorMessage } from './reducer';
|
||||
|
||||
describe('resolveSessionErrorMessage', () => {
|
||||
it('returns a dedicated message for IP session limit errors', () => {
|
||||
@@ -24,162 +19,3 @@ describe('resolveSessionErrorMessage', () => {
|
||||
})).toBe('Backend replica unavailable. Rejoin session.');
|
||||
});
|
||||
});
|
||||
|
||||
function buildContext(overrides: Record<string, unknown> = {}) {
|
||||
const sessionStore = createSessionStore();
|
||||
const rewriteStore = createRewriteStore();
|
||||
const streamStore = createStreamStore();
|
||||
const uiStore = createUiStore();
|
||||
const promptWindowStore = createPromptWindowStore();
|
||||
const avPipeline = {
|
||||
reset: () => {},
|
||||
setStreamCompleted: () => {},
|
||||
noteSegmentInit: () => {},
|
||||
noteSegmentComplete: () => {},
|
||||
maybeStartPlayback: () => {},
|
||||
ensurePipeline: async () => {},
|
||||
};
|
||||
return {
|
||||
sessionStore,
|
||||
promptWindowStore,
|
||||
rewriteStore,
|
||||
streamStore,
|
||||
uiStore,
|
||||
avPipeline,
|
||||
tick: async () => {},
|
||||
defaultAvMime: 'video/mp4',
|
||||
fixedRewriteModel: 'model',
|
||||
parseLatencyMs: () => null,
|
||||
formatPromptWindowEventText: () => '',
|
||||
makePromptId: () => 'generated-id',
|
||||
buildClipLabel: () => 'clip',
|
||||
startSessionCountdown: () => {},
|
||||
clearCountdownInterval: () => {},
|
||||
resetTtffTimer: () => {},
|
||||
startTtffTimer: () => {},
|
||||
preserveArchivedPlaybackSelection: false,
|
||||
finalizeStreamCompletion: async () => {},
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
describe('steering generatingNextScene flow', () => {
|
||||
it('sets generatingNextScene on prompt/sources_resumed in manual mode', async () => {
|
||||
const context = buildContext();
|
||||
await applyNormalizedSocketEvent(
|
||||
{ type: 'prompt/sources_resumed', payload: { segment_idx: 2 } },
|
||||
context,
|
||||
);
|
||||
expect(context.sessionStore.get().generatingNextScene).toBe(true);
|
||||
expect(context.sessionStore.get().waitingForSegmentPrompt).toBe(false);
|
||||
});
|
||||
|
||||
it('does NOT set generatingNextScene on session/auto_extension_updated', async () => {
|
||||
const context = buildContext();
|
||||
await applyNormalizedSocketEvent(
|
||||
{ type: 'session/auto_extension_updated', payload: { enabled: true } },
|
||||
context,
|
||||
);
|
||||
expect(context.sessionStore.get().generatingNextScene).toBe(false);
|
||||
expect(context.sessionStore.get().waitingForSegmentPrompt).toBe(false);
|
||||
});
|
||||
|
||||
it('clears generatingNextScene when segment media arrives', async () => {
|
||||
const context = buildContext();
|
||||
context.sessionStore.patch({ generatingNextScene: true });
|
||||
await applyNormalizedSocketEvent(
|
||||
{
|
||||
type: 'stream/media_init',
|
||||
payload: { segment_idx: 2, stream_id: 's', mime: 'video/mp4' },
|
||||
},
|
||||
context,
|
||||
);
|
||||
expect(context.sessionStore.get().generatingNextScene).toBe(false);
|
||||
});
|
||||
|
||||
it('clears generatingNextScene and returns to waiting on prompt/sources_blocked', async () => {
|
||||
const context = buildContext();
|
||||
context.sessionStore.patch({ generatingNextScene: true });
|
||||
await applyNormalizedSocketEvent(
|
||||
{ type: 'prompt/sources_blocked', payload: { segment_idx: 3 } },
|
||||
context,
|
||||
);
|
||||
expect(context.sessionStore.get().generatingNextScene).toBe(false);
|
||||
expect(context.sessionStore.get().waitingForSegmentPrompt).toBe(true);
|
||||
});
|
||||
|
||||
it('clears generatingNextScene when the opening prompt falls back', async () => {
|
||||
const context = buildContext();
|
||||
context.sessionStore.patch({ generatingNextScene: true });
|
||||
await applyNormalizedSocketEvent(
|
||||
{
|
||||
type: 'prompt/fallback_used',
|
||||
payload: { prompt_id: 'p1', prompt: '', source: 'user_enhancement_failed' },
|
||||
},
|
||||
context,
|
||||
);
|
||||
expect(context.sessionStore.get().generatingNextScene).toBe(false);
|
||||
expect(context.sessionStore.get().waitingForSegmentPrompt).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
describe('opening prompt id tracking', () => {
|
||||
it('routes prompt lifecycle updates to the frontend-recorded opening event', async () => {
|
||||
// The frontend records the opening scene under its own prompt id and sends it
|
||||
// as initial_rollout_prompt_id; the backend echoes it in status updates.
|
||||
const context = buildContext();
|
||||
context.rewriteStore.addPromptEvent({
|
||||
promptId: 'opening-id',
|
||||
status: 'rewrite_requested',
|
||||
source: 'user_rewrite',
|
||||
text: 'a castle at dawn',
|
||||
steeringUserPrompt: true,
|
||||
rawText: 'a castle at dawn',
|
||||
});
|
||||
|
||||
await applyNormalizedSocketEvent(
|
||||
{ type: 'prompt/enhancing', payload: { prompt_id: 'opening-id' } },
|
||||
context,
|
||||
);
|
||||
let opening = (context.rewriteStore.get().promptEvents as Record<string, any>[])
|
||||
.find((e) => e.promptId === 'opening-id');
|
||||
expect(opening?.status).toBe('enhancing');
|
||||
|
||||
await applyNormalizedSocketEvent(
|
||||
{
|
||||
type: 'prompt/fallback_used',
|
||||
payload: { prompt_id: 'opening-id', prompt: '', source: 'user_enhancement_failed' },
|
||||
},
|
||||
context,
|
||||
);
|
||||
opening = (context.rewriteStore.get().promptEvents as Record<string, any>[])
|
||||
.find((e) => e.promptId === 'opening-id');
|
||||
expect(opening?.status).toBe('ready_fallback');
|
||||
// A failed opening is dropped from the steering scene list instead of
|
||||
// lingering as a ghost "Scene 1".
|
||||
expect(opening?.steeringFailed).toBe(true);
|
||||
});
|
||||
|
||||
it('marks a prompt-scoped session/error (e.g. safety block) as steeringFailed', async () => {
|
||||
const context = buildContext();
|
||||
context.rewriteStore.addPromptEvent({
|
||||
promptId: 'blocked-id',
|
||||
status: 'queued',
|
||||
source: 'user_raw',
|
||||
text: 'a blocked prompt',
|
||||
steeringUserPrompt: true,
|
||||
rawText: 'a blocked prompt',
|
||||
});
|
||||
|
||||
await applyNormalizedSocketEvent(
|
||||
{
|
||||
type: 'session/error',
|
||||
payload: { message: 'Prompt blocked by safety filter.', prompt_id: 'blocked-id' },
|
||||
},
|
||||
context,
|
||||
);
|
||||
const blocked = (context.rewriteStore.get().promptEvents as Record<string, any>[])
|
||||
.find((e) => e.promptId === 'blocked-id');
|
||||
expect(blocked?.steeringFailed).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -99,19 +99,7 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
|
||||
status: "ready_fallback",
|
||||
source: payload.source || "user_raw",
|
||||
text: payload.prompt,
|
||||
// Steering: this prompt produced no segment — drop it from the scene list.
|
||||
steeringFailed: true,
|
||||
});
|
||||
// Steering recovery: enhancement failed so the backend enqueued nothing AND won't
|
||||
// re-emit prompt_sources_blocked (its drained flag is still set). Put the user back to
|
||||
// "describe the next scene" ourselves so the generating overlay clears and they can retry.
|
||||
if (!uiStore.get().simpleMode && sessionStore.get().manualContinuationMode) {
|
||||
sessionStore.patch({
|
||||
waitingForSegmentPrompt: true,
|
||||
generatingNextScene: false,
|
||||
sessionNotice: "Couldn't continue from that prompt — try rephrasing the next scene.",
|
||||
});
|
||||
}
|
||||
console.warn("[PromptEnhanceFallback] Prompt extension failed for this request.");
|
||||
return;
|
||||
|
||||
@@ -255,32 +243,19 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
|
||||
}
|
||||
|
||||
case "prompt/sources_blocked":
|
||||
if (sessionStore.get().manualContinuationMode) {
|
||||
sessionStore.patch({ waitingForSegmentPrompt: true, generatingNextScene: false, autoExtensionTimeoutHint: "" });
|
||||
} else {
|
||||
sessionStore.patch({
|
||||
autoExtensionTimeoutHint: uiStore.get().simpleMode ? "" : "blocked on user input, increase prompt count for smoother experience",
|
||||
});
|
||||
}
|
||||
sessionStore.patch({
|
||||
autoExtensionTimeoutHint: uiStore.get().simpleMode ? "" : "blocked on user input, increase prompt count for smoother experience",
|
||||
});
|
||||
return;
|
||||
|
||||
case "prompt/sources_resumed":
|
||||
sessionStore.patch({
|
||||
autoExtensionTimeoutHint: "",
|
||||
waitingForSegmentPrompt: false,
|
||||
// A real prompt was just selected for the next segment; media arriving
|
||||
// (stream/media_init) clears this again.
|
||||
...(sessionStore.get().manualContinuationMode ? { generatingNextScene: true } : {}),
|
||||
});
|
||||
return;
|
||||
|
||||
case "session/auto_extension_updated":
|
||||
// Deliberately does NOT touch generatingNextScene: toggling auto extension
|
||||
// starts no generation.
|
||||
sessionStore.patch({ autoExtensionTimeoutHint: "", waitingForSegmentPrompt: false });
|
||||
console.log("[AutoExtensionUpdated]", {
|
||||
enabled: sessionStore.get().autoExtensionEnabled,
|
||||
});
|
||||
sessionStore.patch({ autoExtensionTimeoutHint: "" });
|
||||
if (event.type === "session/auto_extension_updated") {
|
||||
console.log("[AutoExtensionUpdated]", {
|
||||
enabled: sessionStore.get().autoExtensionEnabled,
|
||||
});
|
||||
}
|
||||
return;
|
||||
|
||||
case "segment/step_complete":
|
||||
@@ -302,7 +277,6 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
|
||||
projectResetPending: false,
|
||||
sessionExpired: true,
|
||||
sessionNotice: "",
|
||||
generatingNextScene: false,
|
||||
});
|
||||
console.log("Session timed out");
|
||||
clearCountdownInterval();
|
||||
@@ -380,8 +354,6 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
|
||||
return;
|
||||
|
||||
case "stream/media_init":
|
||||
// Segment media is arriving — the "Generating next scene" phase is over.
|
||||
sessionStore.patch({ generatingNextScene: false });
|
||||
streamStore.patch({
|
||||
mediaAppendError: null,
|
||||
loadingAnimation: streamStore.get().avPlaybackStarted ? streamStore.get().loadingAnimation : true,
|
||||
@@ -499,30 +471,17 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
|
||||
sessionStore.patch({
|
||||
generationCapReached: false,
|
||||
sessionNotice: "",
|
||||
generatingNextScene: false,
|
||||
});
|
||||
await finalizeStreamCompletion();
|
||||
return;
|
||||
|
||||
case "session/error": {
|
||||
const errorMessage = resolveSessionErrorMessage(payload);
|
||||
if (payload?.prompt_id) {
|
||||
// Prompt-scoped error (e.g. safety-blocked): the prompt produced no
|
||||
// segment, so drop it from the steering scene list.
|
||||
rewriteStore.trackPromptEvent(payload.prompt_id, {
|
||||
steeringFailed: true,
|
||||
});
|
||||
}
|
||||
sessionStore.patch({
|
||||
generationCapReached: false,
|
||||
preservePlaybackOnClose: false,
|
||||
promptExtensionError: "",
|
||||
sessionNotice: errorMessage,
|
||||
// Steering: a blocked/failed prompt produced no segment and the backend won't re-emit
|
||||
// prompt_sources_blocked, so recover the "describe the next scene" state ourselves.
|
||||
...(!uiStore.get().simpleMode && sessionStore.get().manualContinuationMode
|
||||
? { waitingForSegmentPrompt: true, generatingNextScene: false }
|
||||
: {}),
|
||||
});
|
||||
rewriteStore.patch({
|
||||
rewritingSeedPrompts: false,
|
||||
|
||||
@@ -24,9 +24,6 @@ export interface SessionState {
|
||||
livePromptRewriteMode: boolean;
|
||||
sessionExpired: boolean;
|
||||
projectResetPending: boolean;
|
||||
manualContinuationMode: boolean;
|
||||
waitingForSegmentPrompt: boolean;
|
||||
generatingNextScene: boolean;
|
||||
}
|
||||
|
||||
const DEFAULT_SESSION_STATE: SessionState = {
|
||||
@@ -52,12 +49,6 @@ const DEFAULT_SESSION_STATE: SessionState = {
|
||||
livePromptRewriteMode: false,
|
||||
sessionExpired: false,
|
||||
projectResetPending: false,
|
||||
// Steering-only product: every session drives scenes manually. There is no
|
||||
// auto-rollout mode and no UI selector, so this stays true throughout.
|
||||
manualContinuationMode: true,
|
||||
waitingForSegmentPrompt: false,
|
||||
// True from scene submit / prompt selection until the segment's media arrives.
|
||||
generatingNextScene: false,
|
||||
};
|
||||
|
||||
export type SessionStore = ManagedStore<SessionState> & {
|
||||
|
||||
@@ -1,18 +1,5 @@
|
||||
import { createManagedStore, type ManagedStore } from "./createManagedStore";
|
||||
|
||||
// Prompt-history sources driven by the user's own submissions. Curated (non-user)
|
||||
// entries — e.g. a preset's opening scene — feed the steering scene list and are
|
||||
// exempt from the history cap so Scene 1 survives long sessions.
|
||||
export const USER_PROMPT_SOURCES = new Set([
|
||||
"user_raw",
|
||||
"user",
|
||||
"user_enhanced",
|
||||
"user_rewrite",
|
||||
"user_enhancement_failed",
|
||||
]);
|
||||
|
||||
const PROMPT_HISTORY_CAP = 120;
|
||||
|
||||
export interface StreamState {
|
||||
playingSeedPromptIndex: number | null;
|
||||
generatingSeedPromptIndex: number | null;
|
||||
@@ -172,14 +159,10 @@ export function createStreamStore(initialState: Partial<StreamState> = {}): Stre
|
||||
loopIteration: typeof loopIteration === "number" ? loopIteration : null,
|
||||
};
|
||||
|
||||
const nextHistory = [entry, ...state.promptHistory];
|
||||
return {
|
||||
...state,
|
||||
promptHistoryCounter: nextCounter,
|
||||
promptHistory: nextHistory.length > PROMPT_HISTORY_CAP
|
||||
? nextHistory.filter((item, index) =>
|
||||
index < PROMPT_HISTORY_CAP || !USER_PROMPT_SOURCES.has(String(item.source || "")))
|
||||
: nextHistory,
|
||||
promptHistory: [entry, ...state.promptHistory].slice(0, 120),
|
||||
selectedHistoryId: state.selectedHistoryId || (entry.id as string),
|
||||
};
|
||||
});
|
||||
|
||||
@@ -104,11 +104,6 @@ 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`
|
||||
|
||||
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.
|
||||
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`.
|
||||
|
||||
@@ -12,8 +12,7 @@ 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 comparable
|
||||
identity.
|
||||
then compare against the historical baseline for the same model and GPU.
|
||||
|
||||
## Quick start (local)
|
||||
|
||||
@@ -82,9 +81,9 @@ fastvideo/performance/
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
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.
|
||||
`model_id`, `gpu_type`, `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
|
||||
|
||||
@@ -164,11 +163,12 @@ headroom and almost never need touching.
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
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.
|
||||
each available metric, and evaluates the current run with the metric's
|
||||
rolling regression policy. For v2 records, that cohort is `model_id`,
|
||||
`gpu_type`, `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:
|
||||
|
||||
@@ -188,33 +188,10 @@ 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. 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.
|
||||
agent skill to advance the rolling median.
|
||||
|
||||
## Schemas
|
||||
|
||||
@@ -230,14 +207,14 @@ configs and remain loadable. New or migrated configs should use
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3
|
||||
"benchmark_version": 2
|
||||
}
|
||||
```
|
||||
|
||||
`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:
|
||||
`benchmark_id` is still required in this phase because raw artifact names,
|
||||
generated-video directories, normalized record paths, and the current rolling
|
||||
baseline comparator still depend on it. The v2 identity fields are config
|
||||
metadata that make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
@@ -248,26 +225,20 @@ identity fields make the measured workload explicit:
|
||||
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.)
|
||||
migrations do not silently run as v1 configs. Optional v2 metadata fields
|
||||
reserved for follow-up work, such as `metric_threshold_policy` and
|
||||
`quality_metadata`, 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
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, and dashboard cohort grouping land with this change: v2 records
|
||||
compare only within their identity cohort, and a record that opens a NEW
|
||||
cohort is marked `baseline_status: "initialized_new_cohort"` (regression
|
||||
gating starts once that cohort accumulates history). 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.
|
||||
baseline eligible); only static thresholds gate them. Metric-specific
|
||||
threshold policies and promoted baselines remain separate follow-ups.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
@@ -279,7 +250,7 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"benchmark_version": 2,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -318,12 +289,12 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"recipe": {
|
||||
"recipe_schema_version": 2,
|
||||
"recipe_schema_version": 1,
|
||||
"benchmark": {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3
|
||||
"benchmark_version": 2
|
||||
},
|
||||
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
|
||||
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
|
||||
@@ -343,9 +314,6 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"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",
|
||||
@@ -354,12 +322,7 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
}
|
||||
},
|
||||
"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_metadata": { "env": { "IMAGE_VERSION": "py3.12-cuda13.0.0" } },
|
||||
"environment_fingerprint": "env-<sha256-prefix>"
|
||||
}
|
||||
```
|
||||
@@ -375,7 +338,7 @@ result, used as the rolling-baseline source of truth.
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"benchmark_version": 2,
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"commit_sha": "<full sha>",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
@@ -404,10 +367,6 @@ result, used as the rolling-baseline source of truth.
|
||||
"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
|
||||
}
|
||||
```
|
||||
@@ -425,12 +384,9 @@ 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.
|
||||
New records compare only against the same `model_id`, `gpu_type`,
|
||||
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
|
||||
`hardware_profile_id`, and `software_profile_id` cohort.
|
||||
`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
|
||||
@@ -450,16 +406,11 @@ FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
|
||||
| `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` | Performance pytest exit code. Measured static-threshold failures are attributed per record; otherwise a nonzero code reports an infrastructure error. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `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` | `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
|
||||
@@ -474,11 +425,9 @@ 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
|
||||
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after a fixed-threshold
|
||||
failure. Those failed scheduled main runs emit summaries and normalized
|
||||
records, upload records with `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:
|
||||
|
||||
@@ -538,13 +487,10 @@ When the rolling-baseline phase runs, it emits:
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
|
||||
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.
|
||||
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.
|
||||
|
||||
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
|
||||
@@ -564,12 +510,8 @@ When the rolling-baseline phase runs, it emits:
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**`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.
|
||||
**"No baseline for ... Initializing"** — first run for this comparison cohort.
|
||||
Run will pass and (if persisting) seed the first record.
|
||||
|
||||
**Persistent failure right after a torch / kernel / image upgrade** —
|
||||
genuine regression *or* baseline drift. Compare the failing normalized record
|
||||
|
||||
@@ -458,8 +458,6 @@ 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
|
||||
|
||||
@@ -330,7 +330,7 @@ at FastVideo's CI — before the Dynamo-side integration even knows.
|
||||
internal; presets identify them by name on
|
||||
`PipelineSelection.preset`).
|
||||
* `fastvideo.fastvideo_args.FastVideoArgs` (legacy compat type).
|
||||
* `fastvideo.api.compat.*` private helpers
|
||||
* `fastvideo.api.translation.*` private helpers
|
||||
(`_validate_continuation_state` etc.) — the public boundary is
|
||||
`VideoGenerator` + `fastvideo.api`.
|
||||
* Any flat legacy LTX-2 kwarg (`ltx2_refine_upsampler_path`,
|
||||
|
||||
@@ -48,15 +48,17 @@ highest first:
|
||||
`request.model_fields_set` (Pydantic v2). Unset fields do not count,
|
||||
even if the Pydantic model has a schema default for them.
|
||||
2. **`ServeConfig.default_request` (operator-explicit)** — projected via
|
||||
[`explicit_request_updates()`](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/api/compat.py);
|
||||
[`explicit_request_updates()`](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/api/translation.py);
|
||||
only fields the operator actually wrote into the YAML count as
|
||||
defaults. Every other field inherits the schema default rather than
|
||||
being pinned.
|
||||
defaults (an explicit `null` counts as unset). Every other sampling
|
||||
field stays `None` — "inherit the model preset" — and other sections
|
||||
keep their schema defaults without being pinned.
|
||||
3. **Hardcoded fallback** — e.g. `fps = 24`.
|
||||
|
||||
The gate matters: both surfaces carry schema defaults. Without
|
||||
`model_fields_set` / explicit-path tracking, schema defaults would
|
||||
masquerade as intent and silently shadow the other side.
|
||||
The gate matters: the Pydantic surface carries schema defaults and the
|
||||
dataclass surface carries non-None defaults outside `sampling`. Without
|
||||
`model_fields_set` / explicit-path tracking, defaults would masquerade
|
||||
as intent and silently shadow the other side.
|
||||
|
||||
See [`video_api.py::_build_generation_kwargs`](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/entrypoints/openai/video_api.py)
|
||||
for the canonical implementation; the per-request assembly lives there,
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
# 🌊 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).
|
||||
@@ -39,17 +39,20 @@ All you need to generate videos using multi-gpus from state-of-the-art diffusion
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import EngineConfig, GenerationRequest, GeneratorConfig
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
)
|
||||
)
|
||||
|
||||
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.")
|
||||
video = generator.generate_video(prompt)
|
||||
result = generator.generate(GenerationRequest(prompt=prompt))
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OffloadConfig, OutputConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -8,29 +9,32 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
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."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
video = generator.generate(
|
||||
GenerationRequest(prompt=prompt, output=OutputConfig(output_path=OUTPUT_PATH, save_video=True)))
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
@@ -40,7 +44,8 @@ def main():
|
||||
"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.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(prompt=prompt2, output=OutputConfig(output_path=OUTPUT_PATH, save_video=True)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,24 +1,30 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig,
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# image2world example from official repo
|
||||
image_path = "assets/images/bus_terminal.jpg"
|
||||
|
||||
@@ -33,13 +39,16 @@ def main():
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
image_path=str(image_path),
|
||||
num_cond_frames=1,
|
||||
output_path="outputs_video/cosmos2_5_i2w.mp4",
|
||||
save_video=True,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=str(image_path)),
|
||||
output=OutputConfig(
|
||||
output_path="outputs_video/cosmos2_5_i2w.mp4",
|
||||
save_video=True,
|
||||
),
|
||||
extensions={"num_cond_frames": 1},
|
||||
)
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
@@ -47,4 +56,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -1,24 +1,29 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OffloadConfig, OutputConfig,
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Load default sampling parameters (negative_prompt, resolution, steps, etc.)
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
prompt = (
|
||||
"A high-definition video captures the precision of robotic welding in an industrial setting. "
|
||||
"The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. "
|
||||
@@ -34,11 +39,14 @@ def main():
|
||||
"underscoring the ongoing nature of the welding operation."
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
output_path="outputs_video/cosmos2_5_t2w.mp4",
|
||||
save_video=True,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
output_path="outputs_video/cosmos2_5_t2w.mp4",
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
@@ -46,6 +54,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,23 +1,29 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig,
|
||||
OffloadConfig, OutputConfig,
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
# video2world example from official repo
|
||||
video_path = "assets/videos/robot_pouring.mp4"
|
||||
@@ -36,18 +42,19 @@ def main():
|
||||
"The final frame captures the robotic arm with the pitcher finishing the pour, with the glass now filled to a higher level, while the pitcher is slightly tilted but still held securely by the gripper."
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
video_path=str(video_path),
|
||||
num_cond_frames=1,
|
||||
output_path="outputs_video/cosmos2_5_v2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(video_path=str(video_path)),
|
||||
output=OutputConfig(
|
||||
output_path="outputs_video/cosmos2_5_v2w.mp4",
|
||||
save_video=True,
|
||||
),
|
||||
extensions={"num_cond_frames": 1},
|
||||
))
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -2,7 +2,9 @@ import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OffloadConfig, OutputConfig, PipelineSelection,
|
||||
SamplingConfig)
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
def main():
|
||||
@@ -10,30 +12,36 @@ def main():
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
model_name = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(experimental={"VSA_sparsity": 0.8}),
|
||||
))
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. The puddles reflect glowing signs in kanji, advertising ramen, karaoke, and VR arcades. A woman in a translucent raincoat walks briskly with an LED umbrella. Steam rises from a street food cart, and a cat darts across the screen. Raindrops are visible on the camera lens, creating a cinematic bokeh effect."
|
||||
)
|
||||
start_time = time.perf_counter()
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
@@ -46,7 +54,12 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
sampling=SamplingConfig(num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (ComponentConfig, EngineConfig, GenerationRequest,
|
||||
GeneratorConfig, InputConfig, OffloadConfig,
|
||||
OutputConfig, PipelineSelection, SamplingConfig)
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
@@ -16,16 +19,22 @@ def _env_float(name: str, default: float) -> float:
|
||||
|
||||
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",
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(override_pipeline_cls_name="DreamXWorldPipeline"), ),
|
||||
))
|
||||
|
||||
prompt = os.getenv(
|
||||
"DREAMX_WORLD_PROMPT",
|
||||
@@ -37,25 +46,31 @@ def main():
|
||||
"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
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path or None),
|
||||
sampling=SamplingConfig(
|
||||
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),
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
|
||||
),
|
||||
extensions={
|
||||
"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(",")
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
generator.generate_video(prompt, **kwargs)
|
||||
generator.generate(request)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
@@ -1,140 +0,0 @@
|
||||
# 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()
|
||||
@@ -26,6 +26,15 @@ import os
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.models.camera import create_camera_trajectory
|
||||
|
||||
# Model configuration (use GAMECRAFT_MODEL_PATH for local weights)
|
||||
@@ -55,14 +64,20 @@ OUTPUT_PATH = "video_samples_gamecraft"
|
||||
def main():
|
||||
# Initialize generator
|
||||
# FastVideo will automatically download weights from HuggingFace
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_PATH,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=MODEL_PATH,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Video parameters
|
||||
@@ -96,23 +111,27 @@ def main():
|
||||
prompt = DEFAULT_I2V_PROMPT if is_i2v else DEFAULT_PROMPTS["temple"]
|
||||
print(f"Mode: {'I2V' if is_i2v else 'T2V'}, prompt: {prompt[:60]}...")
|
||||
|
||||
gen_kw = dict(
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
camera_states=camera_states,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=50,
|
||||
guidance_scale=6.0,
|
||||
seed=42,
|
||||
fps=24,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
sampling=SamplingConfig(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=50,
|
||||
guidance_scale=6.0,
|
||||
seed=42,
|
||||
fps=24,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
extensions={"camera_states": camera_states},
|
||||
)
|
||||
if is_i2v:
|
||||
gen_kw["image_path"] = image_path
|
||||
generator.generate_video(**gen_kw)
|
||||
request.inputs = InputConfig(image_path=image_path)
|
||||
generator.generate(request)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -22,6 +22,10 @@ Requirements:
|
||||
import argparse
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig,
|
||||
OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
@@ -74,33 +78,47 @@ def main():
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
video = generator.generate_video(
|
||||
args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
image_path=args.image_path,
|
||||
trajectory_type=args.trajectory,
|
||||
movement_distance=args.movement_distance,
|
||||
camera_rotation=args.camera_rotation,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
fps=24,
|
||||
seed=args.seed,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
)
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
inputs=InputConfig(
|
||||
image_path=args.image_path,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
fps=24,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
),
|
||||
extensions={
|
||||
"trajectory_type": args.trajectory,
|
||||
"movement_distance": args.movement_distance,
|
||||
"camera_rotation": args.camera_rotation,
|
||||
},
|
||||
))
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
@@ -1,107 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,6 +1,13 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
@@ -8,17 +15,21 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
generator = VideoGenerator.from_config(GeneratorConfig(
|
||||
model_path="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
@@ -26,7 +37,12 @@ def main():
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
generator.generate(GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(num_frames=81, fps=16),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
@@ -35,8 +51,13 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
generator.generate(GenerationRequest(
|
||||
prompt=prompt2,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(num_frames=81, fps=16),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,12 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
def main():
|
||||
@@ -8,17 +14,23 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
prompt = (
|
||||
@@ -27,7 +39,13 @@ def main():
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
@@ -36,7 +54,13 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
negative_prompt="",
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig,
|
||||
SamplingConfig)
|
||||
from fastvideo.models.dits.hyworld.resolution_utils import get_resolution_from_image
|
||||
|
||||
# Default prompt from HY-WorldPlay run.sh
|
||||
@@ -31,33 +33,45 @@ def main():
|
||||
|
||||
# Initialize generator
|
||||
print("\nInitializing VideoGenerator for HYWorld...")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
image_encoder_cpu_offload=True,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
image_encoder=True,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
# Generate video
|
||||
# The pose string is automatically converted to camera matrices by the pipeline
|
||||
print("\nGenerating video...")
|
||||
generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
image_path=args.image,
|
||||
pose=args.pose, # Camera trajectory control
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
negative_prompt="",
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
height=HEIGHT,
|
||||
width=WIDTH,
|
||||
seed=args.seed,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
inputs=InputConfig(
|
||||
image_path=args.image,
|
||||
pose=args.pose, # Camera trajectory control
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
height=HEIGHT,
|
||||
width=WIDTH,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
),
|
||||
))
|
||||
|
||||
print(f"\nVideo saved to: {args.output_path}")
|
||||
|
||||
|
||||
@@ -1,36 +1,43 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
InputConfig, OffloadConfig, OutputConfig,
|
||||
SamplingConfig)
|
||||
|
||||
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,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="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"
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=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,
|
||||
)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=IMAGE_PATH),
|
||||
sampling=SamplingConfig(height=1024, width=1024, num_frames=121),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,28 +1,41 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OffloadConfig, OutputConfig, SamplingConfig)
|
||||
|
||||
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,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="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"
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=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)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(height=512, width=768, num_frames=121),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
@@ -30,8 +43,13 @@ def main():
|
||||
"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)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
sampling=SamplingConfig(height=512, width=768, num_frames=121),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,24 +1,31 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
num_frames = 81
|
||||
@@ -33,15 +40,23 @@ def main():
|
||||
spatial_scale=8,
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_frames=num_frames,
|
||||
height=480,
|
||||
width=832,
|
||||
c2ws_plucker_emb=c2ws_plucker_emb,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(
|
||||
image_path=image_path,
|
||||
c2ws_plucker_emb=c2ws_plucker_emb,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
num_frames=num_frames,
|
||||
height=480,
|
||||
width=832,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,10 @@ import glob
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
ComponentConfig, EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
InputConfig, OffloadConfig, OutputConfig, PipelineSelection, SamplingConfig,
|
||||
)
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
@@ -45,41 +49,50 @@ SEED = 42
|
||||
def basic_generation():
|
||||
"""
|
||||
Run basic LongCat I2V generation (50 steps at 480p).
|
||||
|
||||
|
||||
This uses the full 50-step denoising process for highest quality.
|
||||
"""
|
||||
print("=" * 60)
|
||||
print("LongCat I2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(experimental={"enable_bsa": False}),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_i2v_basic"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=480, # Square
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(image_path=IMAGE_PATH),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=480, # Square
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(output_path=output_path, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -87,55 +100,70 @@ def basic_generation():
|
||||
def distill_refine_generation():
|
||||
"""
|
||||
Run LongCat I2V with distill+refine pipeline (16 steps + refinement to 768p).
|
||||
|
||||
|
||||
This uses the distilled LoRA for fast 480p generation (16 steps),
|
||||
then refines to 768p using the refinement LoRA with BSA enabled.
|
||||
"""
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat I2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": False,
|
||||
"lora_nickname": "distilled",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_i2v_distill"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=distill_output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=480, # Square
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(image_path=IMAGE_PATH),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=480, # Square
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(output_path=distill_output_path, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 768p)
|
||||
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -143,46 +171,63 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
|
||||
# For BSA [4, 4, 8]: latent must be divisible by 8
|
||||
# 768x768: latent 48x48, 48%8=0 ✓
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=True,
|
||||
bsa_sparsity=0.875,
|
||||
bsa_chunk_q=[4, 4, 4],
|
||||
bsa_chunk_k=[4, 4, 4],
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
refine_generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": True,
|
||||
"bsa_sparsity": 0.875,
|
||||
"bsa_chunk_q": [4, 4, 4],
|
||||
"bsa_chunk_k": [4, 4, 4],
|
||||
"lora_nickname": "refinement",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output_path=refine_output_path,
|
||||
save_video=True,
|
||||
refine_from=distill_video_path,
|
||||
t_thresh=0.5,
|
||||
spatial_refine_only=False,
|
||||
num_cond_frames=0,
|
||||
height=720,
|
||||
width=720,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
refine_generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(refine_from=distill_video_path),
|
||||
sampling=SamplingConfig(
|
||||
height=720,
|
||||
width=720,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(output_path=refine_output_path, save_video=True),
|
||||
extensions={
|
||||
"t_thresh": 0.5,
|
||||
"spatial_refine_only": False,
|
||||
"num_cond_frames": 0,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -192,13 +237,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Image-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
|
||||
@@ -13,6 +13,10 @@ import glob
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
ComponentConfig, EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
InputConfig, OffloadConfig, OutputConfig, PipelineSelection, SamplingConfig,
|
||||
)
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
@@ -38,40 +42,54 @@ SEED = 42
|
||||
def basic_generation():
|
||||
"""
|
||||
Run basic LongCat T2V generation (50 steps at 480p).
|
||||
|
||||
|
||||
This uses the full 50-step denoising process for highest quality.
|
||||
"""
|
||||
print("=" * 60)
|
||||
print("LongCat T2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
experimental={"enable_bsa": False},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_t2v_basic"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -79,54 +97,72 @@ def basic_generation():
|
||||
def distill_refine_generation():
|
||||
"""
|
||||
Run LongCat T2V with distill+refine pipeline (16 steps + refinement to 720p).
|
||||
|
||||
|
||||
This uses the distilled LoRA for fast 480p generation (16 steps),
|
||||
then refines to 720p using the refinement LoRA with BSA enabled.
|
||||
"""
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat T2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": False,
|
||||
"lora_nickname": "distilled",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_t2v_distill"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output_path=distill_output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path=distill_output_path,
|
||||
save_video=True,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 720p)
|
||||
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -134,43 +170,65 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=True,
|
||||
bsa_sparsity=0.875,
|
||||
bsa_chunk_q=[4, 4, 8],
|
||||
bsa_chunk_k=[4, 4, 8],
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
refine_generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": True,
|
||||
"bsa_sparsity": 0.875,
|
||||
"bsa_chunk_q": [4, 4, 8],
|
||||
"bsa_chunk_k": [4, 4, 8],
|
||||
"lora_nickname": "refinement",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output_path=refine_output_path,
|
||||
save_video=True,
|
||||
refine_from=distill_video_path,
|
||||
t_thresh=0.5,
|
||||
spatial_refine_only=False,
|
||||
num_cond_frames=0,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
refine_generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path=refine_output_path,
|
||||
save_video=True,
|
||||
),
|
||||
inputs=InputConfig(
|
||||
refine_from=distill_video_path,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
extensions={
|
||||
"t_thresh": 0.5,
|
||||
"spatial_refine_only": False,
|
||||
"num_cond_frames": 0,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -180,13 +238,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Text-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -194,5 +252,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,10 @@ import glob
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
ComponentConfig, EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig,
|
||||
PipelineSelection, SamplingConfig,
|
||||
)
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
@@ -63,35 +67,49 @@ def basic_generation():
|
||||
"Please provide a valid video path."
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-VC-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
experimental={"enable_bsa": False},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_vc_basic"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
video_path=VIDEO_PATH,
|
||||
num_cond_frames=NUM_COND_FRAMES,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(video_path=VIDEO_PATH),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=50,
|
||||
fps=15,
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
extensions={"num_cond_frames": NUM_COND_FRAMES},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -118,37 +136,55 @@ def distill_refine_generation():
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-VC-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": False,
|
||||
"lora_nickname": "distilled",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_vc_distill"
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
video_path=VIDEO_PATH,
|
||||
num_cond_frames=NUM_COND_FRAMES,
|
||||
output_path=distill_output_path,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(video_path=VIDEO_PATH),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=93,
|
||||
num_inference_steps=16,
|
||||
fps=15,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=distill_output_path,
|
||||
save_video=True,
|
||||
),
|
||||
extensions={"num_cond_frames": NUM_COND_FRAMES},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -166,41 +202,61 @@ def distill_refine_generation():
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=True,
|
||||
bsa_sparsity=0.875,
|
||||
bsa_chunk_q=[4, 4, 8],
|
||||
bsa_chunk_k=[4, 4, 8],
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
refine_generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
),
|
||||
experimental={
|
||||
"enable_bsa": True,
|
||||
"bsa_sparsity": 0.875,
|
||||
"bsa_chunk_q": [4, 4, 8],
|
||||
"bsa_chunk_k": [4, 4, 8],
|
||||
"lora_nickname": "refinement",
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_vc_refine_720p"
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
output_path=refine_output_path,
|
||||
save_video=True,
|
||||
refine_from=distill_video_path,
|
||||
t_thresh=0.5,
|
||||
spatial_refine_only=False,
|
||||
num_cond_frames=0, # For refinement, no conditioning frames
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
|
||||
refine_generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
inputs=InputConfig(refine_from=distill_video_path),
|
||||
sampling=SamplingConfig(
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=50,
|
||||
fps=30,
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=refine_output_path,
|
||||
save_video=True,
|
||||
),
|
||||
extensions={
|
||||
"t_thresh": 0.5,
|
||||
"spatial_refine_only": False,
|
||||
"num_cond_frames": 0, # For refinement, no conditioning frames
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
|
||||
@@ -1,4 +1,11 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
@@ -18,22 +25,32 @@ PROMPT = (
|
||||
|
||||
def main() -> None:
|
||||
# Uses FastVideo default sampling settings for LTX2 base.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Davids048/LTX2-Base-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
),
|
||||
)
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -49,6 +49,11 @@ from pathlib import Path
|
||||
import torch._inductor.config as _inductor
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig, ComponentConfig, EngineConfig, GenerationRequest,
|
||||
GeneratorConfig, OffloadConfig, OutputConfig, PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
@@ -86,9 +91,9 @@ PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
# Per-stage timing helpers --------------------------------------------------
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
def _print_stage_breakdown(result, label: str) -> float | None:
|
||||
"""Print stage execution times and return the sum, or None if missing."""
|
||||
logging_info = result.get("logging_info")
|
||||
logging_info = result.logging_info
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
@@ -104,11 +109,11 @@ def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
result,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
logging_info = result.logging_info
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
@@ -169,34 +174,54 @@ def main() -> None:
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = None
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
# LTX-2.3 distilled uses the two-stage refine pipeline; the refine
|
||||
# LoRA is intentionally empty for the distilled student.
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
# Compile the VAE codec submodules (encoder / decoder) too. The
|
||||
# `LTX2CausalVideoAutoencoder` declares `_compile_conditions` so
|
||||
# `_compile_with_conditions` targets just those submodules and
|
||||
# leaves the surrounding tiling control flow eager — needed for
|
||||
# fullgraph + dynamic=False to succeed. VAE eager decode is
|
||||
# ~1.0s; compiling it brings the stage to ~0.3s.
|
||||
enable_torch_compile_vae=True,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=torch_compile_kwargs,
|
||||
# Keep everything resident — no CPU offload for serving-style runs.
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_root,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
compile=CompileConfig(
|
||||
enabled=True,
|
||||
text_encoder_enabled=True,
|
||||
# Compile the VAE codec submodules (encoder / decoder)
|
||||
# too. The `LTX2CausalVideoAutoencoder` declares
|
||||
# `_compile_conditions` so `_compile_with_conditions`
|
||||
# targets just those submodules and leaves the
|
||||
# surrounding tiling control flow eager — needed for
|
||||
# fullgraph + dynamic=False to succeed. VAE eager decode
|
||||
# is ~1.0s; compiling it brings the stage to ~0.3s.
|
||||
vae_enabled=True,
|
||||
backend=torch_compile_kwargs["backend"],
|
||||
fullgraph=torch_compile_kwargs["fullgraph"],
|
||||
mode=torch_compile_kwargs["mode"],
|
||||
dynamic=torch_compile_kwargs["dynamic"],
|
||||
vae_kwargs=torch_compile_kwargs,
|
||||
),
|
||||
# Keep everything resident — no CPU offload for serving runs.
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
text_encoder=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
vae_tiling=False,
|
||||
# LTX-2.3 distilled uses the two-stage refine pipeline; the
|
||||
# refine LoRA is intentionally empty for the distilled
|
||||
# student.
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
),
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
"enabled": True,
|
||||
"num_inference_steps": 3,
|
||||
"guidance_scale": 1.0,
|
||||
"add_noise": True,
|
||||
}
|
||||
},
|
||||
experimental={"pipeline_config": pipeline_config},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
@@ -206,12 +231,15 @@ def main() -> None:
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
# i2v: anchor the input image at frame 0 with full strength.
|
||||
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
|
||||
# JPEG conditioning image.
|
||||
)
|
||||
|
||||
# i2v: anchor the input image at frame 0 with full strength.
|
||||
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
|
||||
# JPEG conditioning image. These are model-specific knobs routed through
|
||||
# the request extensions escape hatch.
|
||||
common_extensions = dict(
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
warmup_runs = 2
|
||||
@@ -227,10 +255,25 @@ def main() -> None:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=common_kwargs["prompt"],
|
||||
negative_prompt=common_kwargs["negative_prompt"],
|
||||
sampling=SamplingConfig(
|
||||
guidance_scale=common_kwargs["guidance_scale"],
|
||||
height=common_kwargs["height"],
|
||||
width=common_kwargs["width"],
|
||||
num_frames=common_kwargs["num_frames"],
|
||||
fps=common_kwargs["fps"],
|
||||
num_inference_steps=common_kwargs["num_inference_steps"],
|
||||
seed=7,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
save_video=True,
|
||||
),
|
||||
extensions=common_extensions,
|
||||
)
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
@@ -245,19 +288,31 @@ def main() -> None:
|
||||
out_path = OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_run_{m + 1}.mp4"
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
result = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=common_kwargs["prompt"],
|
||||
negative_prompt=common_kwargs["negative_prompt"],
|
||||
sampling=SamplingConfig(
|
||||
guidance_scale=common_kwargs["guidance_scale"],
|
||||
height=common_kwargs["height"],
|
||||
width=common_kwargs["width"],
|
||||
num_frames=common_kwargs["num_frames"],
|
||||
fps=common_kwargs["fps"],
|
||||
num_inference_steps=common_kwargs["num_inference_steps"],
|
||||
seed=2002 + m,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(out_path),
|
||||
save_video=True,
|
||||
),
|
||||
extensions=common_extensions,
|
||||
)
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
e2e = (result.extra.get("e2e_latency") if result is not None else None) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
if result is not None:
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
@@ -17,16 +18,19 @@ import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
num_gpus=4,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/LTX2-Distilled-Diffusers",
|
||||
engine=EngineConfig(num_gpus=4),
|
||||
)
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(output_path=output_path, save_video=True),
|
||||
)
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
@@ -8,6 +8,11 @@ from pathlib import Path
|
||||
import torch
|
||||
import torch._inductor.config
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig, ComponentConfig, EngineConfig, GenerationRequest,
|
||||
GenerationResult, GeneratorConfig, OffloadConfig, OutputConfig,
|
||||
PipelineSelection, SamplingConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
@@ -45,11 +50,11 @@ def load_validation_entries(path: Path) -> list[dict]:
|
||||
|
||||
|
||||
def print_stage_breakdown(
|
||||
result: dict,
|
||||
result: GenerationResult,
|
||||
run_idx: int,
|
||||
num_runs: int,
|
||||
) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
logging_info = result.logging_info
|
||||
if logging_info is None:
|
||||
print(f"[{run_idx}/{num_runs}] Stage breakdown unavailable: no logging_info")
|
||||
return None
|
||||
@@ -70,9 +75,9 @@ def print_stage_breakdown(
|
||||
|
||||
|
||||
def extract_sr_forward_latency(
|
||||
result: dict,
|
||||
result: GenerationResult,
|
||||
) -> tuple[float | None, list[tuple[str, float]], list[str]]:
|
||||
logging_info = result.get("logging_info")
|
||||
logging_info = result.logging_info
|
||||
if logging_info is None:
|
||||
return None, [], []
|
||||
|
||||
@@ -106,11 +111,11 @@ def extract_sr_forward_latency(
|
||||
|
||||
|
||||
def collect_stage_times(
|
||||
result: dict,
|
||||
result: GenerationResult,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
logging_info = result.logging_info
|
||||
if logging_info is None:
|
||||
return
|
||||
stages = getattr(logging_info, "stages", None)
|
||||
@@ -202,26 +207,45 @@ def main() -> None:
|
||||
"dynamic": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
refine_lora_path="", # keep refine LoRA disabled in this repo's typed adapter
|
||||
ltx2_refine_lora_path="", # keep refine LoRA disabled for distilled model
|
||||
ltx2_refine_num_inference_steps=2,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
enable_torch_compile_vae=True,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=torch_compile_kwargs,
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_root,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
text_encoder=False,
|
||||
vae=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=True,
|
||||
text_encoder_enabled=True,
|
||||
vae_enabled=True,
|
||||
backend="inductor",
|
||||
fullgraph=True,
|
||||
dynamic=False,
|
||||
vae_kwargs=torch_compile_kwargs,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
vae_tiling=False,
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
),
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
"enabled": True,
|
||||
"num_inference_steps": 2,
|
||||
"guidance_scale": 1.0,
|
||||
"add_noise": True,
|
||||
}
|
||||
},
|
||||
experimental={
|
||||
"refine_lora_path": "", # keep refine LoRA disabled in this repo's typed adapter
|
||||
"pipeline_config": pipeline_config,
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
run_times: list[float] = []
|
||||
@@ -243,25 +267,31 @@ def main() -> None:
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
prompt=prompt,
|
||||
output_path=str(output_path),
|
||||
fps=24,
|
||||
seed=10,
|
||||
save_video=True,
|
||||
guidance_scale=1.0,
|
||||
height=benchmark_entry.get("height", 1088),
|
||||
width=benchmark_entry.get("width", 1920),
|
||||
num_frames=121,
|
||||
num_inference_steps=5,
|
||||
# image_path="examples/inference/basic/prompt1.png",
|
||||
# ltx2_image_crf=0.0
|
||||
result = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
fps=24,
|
||||
seed=10,
|
||||
guidance_scale=1.0,
|
||||
height=benchmark_entry.get("height", 1088),
|
||||
width=benchmark_entry.get("width", 1920),
|
||||
num_frames=121,
|
||||
num_inference_steps=5,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_path),
|
||||
save_video=True,
|
||||
),
|
||||
# inputs=InputConfig(image_path="examples/inference/basic/prompt1.png"),
|
||||
# extensions={"ltx2_image_crf": 0.0},
|
||||
)
|
||||
)
|
||||
if os.environ.get("FASTVIDEO_STAGE_LOGGING") == "0":
|
||||
torch.cuda.synchronize()
|
||||
|
||||
elapsed = result.get("generation_time") if isinstance(result, dict) else None
|
||||
e2e_elapsed = result.get("e2e_latency") if isinstance(result, dict) else None
|
||||
elapsed = result.generation_time if isinstance(result, GenerationResult) else None
|
||||
e2e_elapsed = result.extra.get("e2e_latency") if isinstance(result, GenerationResult) else None
|
||||
if elapsed is None:
|
||||
elapsed = time.perf_counter() - start
|
||||
if e2e_elapsed is None:
|
||||
@@ -272,7 +302,7 @@ def main() -> None:
|
||||
print(f"[{i + 1}/{num_runs}] Generation time: {elapsed:.2f}s")
|
||||
print(f"[{i + 1}/{num_runs}] End-to-end latency: {e2e_elapsed:.2f}s")
|
||||
|
||||
if isinstance(result, dict):
|
||||
if isinstance(result, GenerationResult):
|
||||
stage_sum = print_stage_breakdown(result, i + 1, num_runs)
|
||||
if stage_sum is not None:
|
||||
non_stage_overhead = e2e_elapsed - stage_sum
|
||||
|
||||
@@ -1,18 +1,27 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig,
|
||||
OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_lucy_edit"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"decart-ai/Lucy-Edit-Dev",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="decart-ai/Lucy-Edit-Dev",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
prompt = ("Change the apron and blouse to a classic clown costume: satin "
|
||||
"polka-dot jumpsuit in bright primary colors, ruffled white collar, "
|
||||
@@ -20,18 +29,20 @@ def main():
|
||||
"foam nose; soft window light from left, eye-level medium shot.")
|
||||
video_path = "https://d2drjpuinn46lb.cloudfront.net/painter_original_edit.mp4"
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
negative_prompt="",
|
||||
video_path=video_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
fps=24,
|
||||
guidance_scale=5.0,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
inputs=InputConfig(video_path=video_path),
|
||||
sampling=SamplingConfig(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
fps=24,
|
||||
guidance_scale=5.0,
|
||||
),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig,
|
||||
SamplingConfig)
|
||||
from fastvideo.models.dits.matrixgame2.utils import create_action_presets
|
||||
|
||||
import torch
|
||||
@@ -38,35 +40,48 @@ def main():
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
num_frames = 597
|
||||
actions = create_action_presets(num_frames, keyboard_dim=config["keyboard_dim"])
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
|
||||
generator.generate_video(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt="",
|
||||
inputs=InputConfig(
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrixgame2.utils import get_current_action_async, expand_action_to_frames
|
||||
from fastvideo.api import EngineConfig, GeneratorConfig, OffloadConfig
|
||||
|
||||
import torch
|
||||
import asyncio
|
||||
@@ -42,17 +43,23 @@ async def main():
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = StreamingVideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = StreamingVideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
max_blocks = 50
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
MODEL_PATH = "FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers"
|
||||
IMAGE_URL = "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-3/demo_images/001/image.png"
|
||||
@@ -7,28 +10,38 @@ OUTPUT_PATH = "video_samples_matrixgame3"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_PATH,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=MODEL_PATH,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
))
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_URL,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_frames=57,
|
||||
num_inference_steps=3,
|
||||
guidance_scale=1.0,
|
||||
seed=42,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
inputs=InputConfig(image_path=IMAGE_URL),
|
||||
sampling=SamplingConfig(
|
||||
height=720,
|
||||
width=1280,
|
||||
num_frames=57,
|
||||
num_inference_steps=3,
|
||||
guidance_scale=1.0,
|
||||
seed=42,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,40 +1,56 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
disable_autocast=False,
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
experimental={"pipeline_config": config},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
sampling_param.num_frames = 25 # Reduce from default 81 to 25 frames bc we have to use the SDPA attn backend for mps
|
||||
sampling_param.height = 256
|
||||
sampling_param.width = 256
|
||||
# Reduce from default 81 to 25 frames bc we have to use the SDPA attn backend for mps
|
||||
sampling = SamplingConfig(
|
||||
num_frames=25,
|
||||
height=256,
|
||||
width=256,
|
||||
)
|
||||
|
||||
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.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
video = generator.generate(GenerationRequest(prompt=prompt, sampling=sampling))
|
||||
|
||||
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.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
video2 = generator.generate(GenerationRequest(prompt=prompt2, sampling=sampling))
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OffloadConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -8,17 +10,23 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
distributed_executor_backend="ray",
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
execution_backend="ray",
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
@@ -27,7 +35,8 @@ def main():
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
video = generator.generate(
|
||||
GenerationRequest(prompt=prompt, output=OutputConfig(output_path=OUTPUT_PATH, save_video=True)))
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
@@ -37,7 +46,8 @@ def main():
|
||||
"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.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(prompt=prompt2, output=OutputConfig(output_path=OUTPUT_PATH, save_video=True)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -85,24 +85,35 @@ def main() -> None:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OffloadConfig, OutputConfig,
|
||||
ParallelismConfig, PipelineSelection, SamplingConfig,
|
||||
)
|
||||
|
||||
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)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=False,
|
||||
parallelism=ParallelismConfig(
|
||||
sp_size=1,
|
||||
tp_size=1,
|
||||
),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=False,
|
||||
vae=False,
|
||||
image_encoder=False,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
@@ -113,20 +124,25 @@ def main() -> None:
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[sd35] 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,
|
||||
"seed": seed,
|
||||
"negative_prompt": args.negative,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt=args.negative,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance,
|
||||
seed=seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
print(f"[sd35] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_causal"
|
||||
def main():
|
||||
@@ -9,23 +17,33 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
text_encoder=False,
|
||||
dit=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
|
||||
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."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
video = generator.generate(request)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,8 +1,17 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
|
||||
def main():
|
||||
@@ -10,26 +19,37 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
dit_precision="fp32",
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
experimental={
|
||||
"dit_precision": "fp32",
|
||||
"dmd_denoising_steps": [1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
|
||||
sampling_param.num_frames = 81
|
||||
sampling_param.width = 832
|
||||
sampling_param.height = 480
|
||||
sampling_param.seed = 1000
|
||||
sampling = SamplingConfig(
|
||||
num_frames=81,
|
||||
width=832,
|
||||
height=480,
|
||||
seed=1000,
|
||||
)
|
||||
|
||||
with open("assets/prompts/mixkit_i2v.jsonl", "r") as f:
|
||||
prompt_image_pairs = json.load(f)
|
||||
@@ -37,7 +57,14 @@ def main():
|
||||
for prompt_image_pair in prompt_image_pairs:
|
||||
prompt = prompt_image_pair["prompt"]
|
||||
image_path = prompt_image_pair["image_path"]
|
||||
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=sampling,
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
ComponentConfig, EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OffloadConfig, OutputConfig, PipelineSelection, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
def main():
|
||||
@@ -10,34 +12,49 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
num_frame_per_block=7,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
transformer_weights="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
transformer_2_weights="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
),
|
||||
experimental={
|
||||
"dmd_denoising_steps": [1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
"num_frame_per_block": 7,
|
||||
},
|
||||
),
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
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, num_frames=81)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -51,25 +51,31 @@ Prerequisites:
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OutputConfig)
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
))
|
||||
output_path = "outputs_audio/stable_audio_basic/output_stable_audio.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# 6-second clip; the model max is ~47.5s.
|
||||
audio_end_in_s=6.0,
|
||||
# The registered preset gives 100 steps + CFG=7.0 by default;
|
||||
# override num_inference_steps / guidance_scale here for QA.
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
# 6-second clip; the model max is ~47.5s.
|
||||
extensions={"audio_end_in_s": 6.0},
|
||||
# The registered preset gives 100 steps + CFG=7.0 by default;
|
||||
# override num_inference_steps / guidance_scale here for QA.
|
||||
))
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
|
||||
@@ -48,6 +48,12 @@ Picking `init_audio_strength` (0.0 to 1.0):
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
|
||||
PROMPT = "Change the piano to a cello playing the same notes"
|
||||
# Path to any audio-bearing file (wav, mp3, mp4, m4a, flac, ...).
|
||||
@@ -58,18 +64,24 @@ INIT_AUDIO_STRENGTH = 0.6
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=6.0,
|
||||
init_audio=INIT_AUDIO_PATH,
|
||||
init_audio_strength=INIT_AUDIO_STRENGTH,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
))
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
|
||||
save_video=True,
|
||||
),
|
||||
extensions={
|
||||
"audio_end_in_s": 6.0,
|
||||
"init_audio": INIT_AUDIO_PATH,
|
||||
"init_audio_strength": INIT_AUDIO_STRENGTH,
|
||||
},
|
||||
))
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
|
||||
@@ -48,6 +48,9 @@ Prerequisites: same as `basic_stable_audio.py`.
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GeneratorConfig, GenerationRequest, OutputConfig,
|
||||
)
|
||||
|
||||
PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
|
||||
# Required: path to the reference audio file (wav, mp3, mp4, m4a, flac,
|
||||
@@ -64,19 +67,25 @@ def main() -> None:
|
||||
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=TOTAL_SECONDS,
|
||||
inpaint_audio=REFERENCE_AUDIO_PATH,
|
||||
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
|
||||
inpaint_mask=(KEEP_SECONDS, TOTAL_SECONDS),
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
))
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(
|
||||
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
|
||||
save_video=True,
|
||||
),
|
||||
extensions={
|
||||
"audio_end_in_s": TOTAL_SECONDS,
|
||||
"inpaint_audio": REFERENCE_AUDIO_PATH,
|
||||
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
|
||||
"inpaint_mask": (KEEP_SECONDS, TOTAL_SECONDS),
|
||||
},
|
||||
))
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
|
||||
@@ -28,24 +28,27 @@ Prerequisites: same as `basic_stable_audio.py`. The converted repo is
|
||||
public so no gated-access flow is required.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OutputConfig)
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-small-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="FastVideo/stable-audio-open-small-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
))
|
||||
output_path = "outputs_audio/stable_audio_small/output_stable_audio_small.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
|
||||
# at or below that.
|
||||
audio_end_in_s=6.0,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(output_path=output_path, save_video=True),
|
||||
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
|
||||
# at or below that.
|
||||
extensions={"audio_end_in_s": 6.0},
|
||||
))
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,9 @@ import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion"
|
||||
|
||||
@@ -11,14 +14,17 @@ OUTPUT_PATH = "video_samples_turbodiffusion"
|
||||
def main() -> None:
|
||||
# TurboDiffusion: 1-4 step video generation using RCM scheduler + SLA attention
|
||||
# FastVideo will automatically use TurboDiffusionPipeline when specified
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
),
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
)
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
@@ -28,11 +34,17 @@ def main() -> None:
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
seed=42,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
@@ -43,11 +55,17 @@ def main() -> None:
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
sampling=SamplingConfig(
|
||||
seed=42,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,9 @@ import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_turbodiffusion_14B"
|
||||
|
||||
@@ -11,10 +14,12 @@ OUTPUT_PATH = "video_samples_turbodiffusion_14B"
|
||||
def main() -> None:
|
||||
# TurboDiffusion 14B: 1-4 step video generation using RCM scheduler + SLA attention
|
||||
# FastVideo will automatically use TurboDiffusionPipeline when specified
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers",
|
||||
# 14B model needs more GPUs
|
||||
num_gpus=2,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="loayrashid/TurboWan2.1-T2V-14B-Diffusers",
|
||||
# 14B model needs more GPUs
|
||||
engine=EngineConfig(num_gpus=2),
|
||||
)
|
||||
)
|
||||
|
||||
prompt = (
|
||||
@@ -22,11 +27,12 @@ def main() -> None:
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
sampling=SamplingConfig(seed=42),
|
||||
)
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
@@ -37,11 +43,12 @@ def main() -> None:
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
sampling=SamplingConfig(seed=42),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,10 @@ import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLA_ATTN"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig,
|
||||
OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
# Use local model path
|
||||
MODEL_PATH = "loayrashid/TurboWan2.2-I2V-A14B-Diffusers"
|
||||
@@ -12,9 +16,11 @@ OUTPUT_PATH = "video_samples_turbodiffusion_i2v"
|
||||
|
||||
def main() -> None:
|
||||
# TurboDiffusion I2V: 1-4 step image-to-video generation
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_PATH,
|
||||
num_gpus=2,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=MODEL_PATH,
|
||||
engine=EngineConfig(num_gpus=2),
|
||||
)
|
||||
)
|
||||
|
||||
# Example prompt and image for I2V
|
||||
@@ -24,12 +30,13 @@ def main() -> None:
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
seed=42,
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=SamplingConfig(seed=42),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OffloadConfig,
|
||||
OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
|
||||
def main():
|
||||
@@ -8,30 +10,37 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
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=720, width=1280, num_frames=81)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(height=720, width=1280, num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
@@ -41,8 +50,14 @@ def main():
|
||||
"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=720, width=1280, num_frames=81)
|
||||
_ = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
sampling=SamplingConfig(height=720, width=1280, num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,12 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_1_Fun"
|
||||
OUTPUT_NAME = "wan2.1_test"
|
||||
@@ -9,18 +15,24 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
|
||||
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
|
||||
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
prompt = "一位年轻女性穿着一件粉色的连衣裙,裙子上有白色的装饰和粉色的纽扣。她的头发是紫色的,头上戴着一个红色的大蝴蝶结,显得非常可爱和精致。她还戴着一个红色的领结,整体造型充满了少女感和活力。她的表情温柔,双手轻轻交叉放在身前,姿态优雅。背景是简单的灰色,没有任何多余的装饰,使得人物更加突出。她的妆容清淡自然,突显了她的清新气质。整体画面给人一种甜美、梦幻的感觉,仿佛置身于童话世界中。"
|
||||
@@ -30,7 +42,14 @@ def main():
|
||||
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
|
||||
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
|
||||
|
||||
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
inputs=InputConfig(image_path=image_path, video_path=control_video_path),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig,
|
||||
OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
|
||||
def main():
|
||||
@@ -8,23 +10,36 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True, # DiT need to be offloaded for MoE
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, height=832, width=480, num_frames=81)
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=SamplingConfig(height=832, width=480, num_frames=81),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, InputConfig, OffloadConfig, OutputConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
def main():
|
||||
@@ -7,22 +10,34 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
offload=OffloadConfig(
|
||||
dit=True,
|
||||
vae=False,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# I2V is triggered just by passing in an image_path argument
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path)
|
||||
video = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
@@ -34,8 +49,13 @@ def main():
|
||||
"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.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,120 +0,0 @@
|
||||
# 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()
|
||||
@@ -50,23 +50,33 @@ N_DUP = 4 # how many times to duplicate the video for the gen/ref corpora
|
||||
def generate_one_ltx2_video() -> str:
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OutputConfig, SamplingConfig)
|
||||
|
||||
Path(OUTPUT_PATH).parent.mkdir(parents=True, exist_ok=True)
|
||||
# Davids048/LTX2-Base-Diffusers is the audio-capable LTX-2 checkpoint
|
||||
# (the Distilled variant ships without the audio VAE, so its mp4
|
||||
# audio track is silence/noise — unusable for audio.* metrics).
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Davids048/LTX2-Base-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
)
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_frames=121, # ~5s @ 24 fps — long enough for audio.desync (Synchformer ≥14 segments)
|
||||
height=480,
|
||||
width=832,
|
||||
fps=24,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=121, # ~5s @ 24 fps — long enough for audio.desync (Synchformer ≥14 segments)
|
||||
height=480,
|
||||
width=832,
|
||||
fps=24,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
generator.shutdown()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -21,6 +21,10 @@ Install: ``uv pip install -e .[eval-audio]`` covers both metrics here
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
PROMPT = (
|
||||
@@ -39,20 +43,26 @@ METRICS = [
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Davids048/LTX2-Base-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
))
|
||||
|
||||
output_path = "outputs_video/ltx2_audio_eval/output.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
)
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
),
|
||||
))
|
||||
generator.shutdown()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -22,6 +22,10 @@ sharing, or run on a smaller-resolution generation.
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.eval import Evaluator
|
||||
from fastvideo.eval.io import build_eval_kwargs
|
||||
|
||||
@@ -58,19 +62,20 @@ METRICS = [
|
||||
|
||||
def main() -> None:
|
||||
# ----- generation (matches examples/inference/basic/basic_ltx2.py) -----
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path="Davids048/LTX2-Base-Diffusers",
|
||||
engine=EngineConfig(num_gpus=1),
|
||||
)
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
output=OutputConfig(output_path=output_path, save_video=True),
|
||||
sampling=SamplingConfig(num_frames=121, height=1088, width=1920),
|
||||
)
|
||||
)
|
||||
generator.shutdown()
|
||||
# Free residual CUDA memory the generator left behind so the
|
||||
|
||||
@@ -45,6 +45,9 @@ def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig, GenerationRequest, GeneratorConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(row, videos_dir / _expected_filename(row)) for row in rows]
|
||||
@@ -55,13 +58,16 @@ def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(rows)} scenarios to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
gen = VideoGenerator.from_config(GeneratorConfig(
|
||||
model_path=model, engine=EngineConfig(num_gpus=num_gpus),
|
||||
))
|
||||
try:
|
||||
for row, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=row["prompt"], output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
gen.generate(GenerationRequest(
|
||||
prompt=row["prompt"],
|
||||
sampling=SamplingConfig(num_frames=num_frames, height=height, width=width),
|
||||
output=OutputConfig(output_path=str(out_path), save_video=True),
|
||||
))
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
@@ -43,6 +43,8 @@ def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (EngineConfig, GenerationRequest, GeneratorConfig,
|
||||
OutputConfig, SamplingConfig)
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(p, videos_dir / f"{_slugify(p)}.mp4") for p in prompts]
|
||||
@@ -53,13 +55,15 @@ def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(prompts)} prompts to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
gen = VideoGenerator.from_config(GeneratorConfig(
|
||||
model_path=model, engine=EngineConfig(num_gpus=num_gpus)))
|
||||
try:
|
||||
for prompt, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=prompt, output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
gen.generate(GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(num_frames=num_frames, height=height, width=width),
|
||||
output=OutputConfig(output_path=str(out_path), save_video=True),
|
||||
))
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
@@ -33,6 +33,13 @@ import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
@@ -99,16 +106,27 @@ def generate(args: argparse.Namespace) -> Path:
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
print(f"[gen] loading {args.model} ({args.num_gpus} GPU)...")
|
||||
generator = VideoGenerator.from_pretrained(args.model, num_gpus=args.num_gpus)
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model,
|
||||
engine=EngineConfig(num_gpus=args.num_gpus),
|
||||
)
|
||||
)
|
||||
try:
|
||||
print(f"[gen] generating to {out}...")
|
||||
generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
output_path=str(out),
|
||||
save_video=True,
|
||||
num_frames=args.num_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=args.num_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(out),
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
@@ -33,7 +33,7 @@ This demo initializes a `VideoGenerator` with the minimum required arguments for
|
||||
|
||||
The core functionality is in the `generate_video` function, which:
|
||||
1. Processes user inputs
|
||||
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
|
||||
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate(GenerationRequest(...))`)
|
||||
|
||||
## Gradio Interface
|
||||
|
||||
|
||||
@@ -5,7 +5,13 @@ import time
|
||||
|
||||
import gradio as gr
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api import (
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
SamplingParam,
|
||||
)
|
||||
from copy import deepcopy
|
||||
|
||||
|
||||
@@ -129,9 +135,22 @@ def create_gradio_interface(default_params: dict[str, SamplingParam], generators
|
||||
output_dir = "outputs/"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
start_time = time.time()
|
||||
result = generator.generate_video(prompt=prompt, sampling_param=params, save_video=True, return_frames=False)
|
||||
result = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt=params.negative_prompt,
|
||||
sampling=SamplingConfig(
|
||||
seed=int(params.seed),
|
||||
guidance_scale=params.guidance_scale,
|
||||
num_frames=int(params.num_frames),
|
||||
height=int(params.height),
|
||||
width=int(params.width),
|
||||
),
|
||||
output=OutputConfig(save_video=True, return_frames=False),
|
||||
)
|
||||
)
|
||||
inference_time = time.time() - start_time
|
||||
logging_info = result.get("logging_info", None)
|
||||
logging_info = result.logging_info
|
||||
if logging_info:
|
||||
stage_names = logging_info.get_execution_order()
|
||||
stage_execution_times = [
|
||||
@@ -550,7 +569,7 @@ def main():
|
||||
for model_path in model_paths:
|
||||
print(f"Loading model: {model_path}")
|
||||
setup_model_environment(model_path)
|
||||
generators[model_path] = VideoGenerator.from_pretrained(model_path)
|
||||
generators[model_path] = VideoGenerator.from_config(GeneratorConfig(model_path=model_path))
|
||||
default_params[model_path] = SamplingParam.from_pretrained(model_path)
|
||||
demo = create_gradio_interface(default_params, generators)
|
||||
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
|
||||
|
||||
@@ -55,10 +55,11 @@ demo can actually boot:
|
||||
`fastvideo/fastvideo_args.py` currently wires only `ltx2_vae_tiling`.
|
||||
The backing stages (`ltx2_refine.py`, `ltx2_i2v_conditioning.py`) are
|
||||
also missing from `fastvideo/pipelines/stages/`.
|
||||
3. **`fastvideo.configs.sample.base.SamplingParam`** — the import path used
|
||||
by this demo. Upstream moved sampling params to
|
||||
`fastvideo.api.sampling_param`. A re-export shim at the old path, or an
|
||||
import update here once the other two prereqs land, will resolve it.
|
||||
3. **`SamplingParam`** — now imported from `fastvideo.api` (the public
|
||||
re-export of `fastvideo.api.sampling_param`); the old
|
||||
`fastvideo.configs.sample.base` path was removed upstream. `SamplingParam`
|
||||
here only sources model-default slider values — generation itself runs
|
||||
through the typed `GenerationRequest` / `generator.generate(...)` path.
|
||||
|
||||
## Environment variables
|
||||
|
||||
|
||||
@@ -4,8 +4,16 @@ from pathlib import Path
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
ComponentConfig,
|
||||
EngineConfig,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
PipelineSelection,
|
||||
SamplingParam,
|
||||
)
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.layers.quantization.fp4_config import FP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
@@ -48,28 +56,44 @@ def main():
|
||||
refine_upsampler_path = resolve_refine_upsampler_path(resolved_model_path)
|
||||
print(f"Using refine upsampler: {refine_upsampler_path}")
|
||||
|
||||
generators[model_path] = VideoGenerator.from_pretrained(
|
||||
str(resolved_model_path),
|
||||
num_gpus=1,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="", # disable refine LoRA for distilled model
|
||||
ltx2_refine_num_inference_steps=2,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
torch_compile_kwargs={
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "max-autotune-no-cudagraphs",
|
||||
"dynamic": False,
|
||||
},
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
generators[model_path] = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=str(resolved_model_path),
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=False,
|
||||
text_encoder=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=True,
|
||||
text_encoder_enabled=True,
|
||||
backend="inductor",
|
||||
fullgraph=True,
|
||||
mode="max-autotune-no-cudagraphs",
|
||||
dynamic=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
# Empty refine LoRA path (distilled needs none) -> omit.
|
||||
),
|
||||
vae_tiling=False,
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
"enabled": True,
|
||||
"num_inference_steps": 2,
|
||||
"guidance_scale": 1.0,
|
||||
"add_noise": True,
|
||||
},
|
||||
},
|
||||
# PipelineConfig object (with FP4 quant wired on above) has
|
||||
# no first-class typed field; route via experimental.
|
||||
experimental={"pipeline_config": pipeline_config},
|
||||
),
|
||||
)
|
||||
)
|
||||
default_params[model_path] = apply_ltx2_defaults(
|
||||
SamplingParam.from_pretrained(str(resolved_model_path))
|
||||
|
||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
||||
import torch
|
||||
import torch._inductor.config
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api import SamplingParam
|
||||
|
||||
LOCAL_DEMO_DIR = Path(__file__).resolve().parent
|
||||
CLASSIFIER_DIR = Path(
|
||||
|
||||
@@ -5,8 +5,14 @@ from copy import deepcopy
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
SamplingParam,
|
||||
)
|
||||
|
||||
from .config import (
|
||||
DEFAULT_FPS,
|
||||
@@ -69,40 +75,38 @@ def create_gradio_interface(default_params: dict[str, SamplingParam], generators
|
||||
output_path = str(OUTPUT_DIR / video_filename)
|
||||
params.output_path = output_path
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
prompt=prompt,
|
||||
output_path=output_path,
|
||||
fps=DEFAULT_FPS,
|
||||
seed=int(params.seed),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
guidance_scale=float(params.guidance_scale),
|
||||
height=int(params.height),
|
||||
width=int(params.width),
|
||||
num_frames=int(params.num_frames),
|
||||
num_inference_steps=DEFAULT_NUM_INFERENCE_STEPS,
|
||||
negative_prompt=params.negative_prompt,
|
||||
image_path=params.image_path,
|
||||
ltx2_image_crf=0.0
|
||||
result = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt=params.negative_prompt,
|
||||
inputs=InputConfig(image_path=params.image_path),
|
||||
sampling=SamplingConfig(
|
||||
seed=int(params.seed),
|
||||
fps=DEFAULT_FPS,
|
||||
guidance_scale=float(params.guidance_scale),
|
||||
height=int(params.height),
|
||||
width=int(params.width),
|
||||
num_frames=int(params.num_frames),
|
||||
num_inference_steps=DEFAULT_NUM_INFERENCE_STEPS,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
# LTX-2 i2v knob without a first-class typed field yet.
|
||||
extensions={"ltx2_image_crf": 0.0},
|
||||
)
|
||||
)
|
||||
wall_time = time.perf_counter() - start_time
|
||||
generation_time = (
|
||||
result.get("generation_time")
|
||||
if isinstance(result, dict) else None
|
||||
)
|
||||
e2e_latency = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
)
|
||||
generation_time = result.generation_time
|
||||
e2e_latency = result.extra.get("e2e_latency")
|
||||
if generation_time is None:
|
||||
generation_time = wall_time
|
||||
if e2e_latency is None:
|
||||
e2e_latency = wall_time
|
||||
resolved_output_path = (
|
||||
result.get("output_path", output_path)
|
||||
if isinstance(result, dict) else output_path
|
||||
)
|
||||
logging_info = result.get("logging_info", None) if isinstance(result, dict) else None
|
||||
resolved_output_path = result.video_path or output_path
|
||||
logging_info = result.logging_info
|
||||
if logging_info:
|
||||
stage_names = logging_info.get_execution_order()
|
||||
stage_execution_times = [
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user