Compare commits

...
68 Commits
Author SHA1 Message Date
SolitaryThinker 59cf9d6fa2 [fix] dreamverse: steering session-flow, mock parity, and init-image validation
Steering protocol/UI fixes that interlock via one contract change — the
frontend now sends initial_rollout_prompt_id (its own prompt id for the
opening scene) in session-init and reset payloads, and both backend pre-seed
sites use it, so prompt lifecycle events (enhancing/ready/fallback) finally
match the frontend's records instead of no-oping against a backend uuid.

- controller: suppress prompt_sources_blocked while a submission is queued or
  enhancing (fixes the blank player at typed-prompt session start); segment
  prompt logging is now opt-in via DREAMVERSE_SEGMENT_PROMPT_LOG, written off
  the event loop, warning on failure (drops the hardcoded author-machine path
  and bare except).
- web: 'Generating next scene' overlay driven by explicit generation state
  instead of waitingForSegmentPrompt prop edges, so auto_extension_updated
  while idle can't strand it over the video; steering scene history is no
  longer truncated by the 24-event prompt feed cap and keeps stable numbers
  for 30+ scene sessions.
- mock server: manual_continuation_mode parity (no rewrite-flow wait, honors
  initial_rollout_prompt_id, no segment cap in manual mode) so the GPU-less
  dev backend works with the steering-only frontend.
- init image: client-side type/size validation on picker and paste paths
  (png/jpeg/webp, 15MB) with a visible error, and ws_max_size raised to 32MiB
  on both servers so a legitimate ~15MB image reaches the backend's own
  validation instead of tripping uvicorn's 16MiB frame cap.
2026-07-16 20:39:32 -07:00
SolitaryThinker 0c90c328c5 [fix] ltx2: drop unreachable video_position_offset_sec kwargs fallback
video_position_offset_sec is a declared parameter of forward, so a caller's
keyword binds to it and never lands in **kwargs — the added
kwargs.get("video_position_offset_sec", 0.0) block was unreachable dead code
with a false comment, and it rebound the local to 0.0. The pre-existing
application of the offset is the only live path; behavior is unchanged.
2026-07-16 20:39:32 -07:00
SolitaryThinker cfb54a3b2a [fix] dreamverse: env-tunable session timeout, restore warmup watchdog, derive SP size in launch script
- SESSION_TIMEOUT_SECONDS: default back to 300 and now reads the
  DREAMVERSE_SESSION_TIMEOUT_SECONDS env var that launch-dreamverse.sh was
  already exporting (previously nothing read it, and the hardcoded 1800 made
  every deployment hold idle GPU-pool slots 6x longer). Fixes the stale
  five-minute-timeout test and adds an override test.
- STARTUP_WARMUP_TIMEOUT_SECONDS: default back to 2400 so the warmup watchdog
  works again; launch-dreamverse.sh exports the existing
  FASTVIDEO_STARTUP_WARMUP_TIMEOUT_SECONDS override (24000) for the slow
  GB200 max-autotune boot.
- launch-dreamverse.sh: derive the DREAMVERSE_SP_SIZE default from the number
  of visible GPUs so the documented CUDA_VISIBLE_DEVICES=0 invocation no
  longer crashes gpu_pool with 'Not enough GPUs'; explicit env still wins.
  setup-dreamverse-env.sh's printed instruction now matches. gpu_pool raises
  a friendly error for non-integer CUDA_VISIBLE_DEVICES entries.
- video_generation.py: drop unused ParallelismConfig import (F401).
2026-07-16 20:39:13 -07:00
SolitaryThinker deb51f1dcc [fix] dreamverse: return last outer JSON object when parsing enhancer replies
The free-form scanner attempted raw_decode at every '{', including braces
inside an already-decoded object, so 'return the last decodable object'
(deliberate: chain-of-thought models emit draft JSON before the final answer)
actually returned the innermost/trailing nested fragment — e.g.
{"next_prompt": ..., "style": {"mood": "noir"}} parsed to {"mood": "noir"}
and enhancement fell back, failing steering requests.

Scan with a position cursor instead: skip past the consumed span of every
successfully decoded object, and skip malformed/truncated spans wholesale via
a balanced-brace scan so their nested fragments can't displace an earlier
complete object. Regression tests cover nested values, sequential drafts,
truncated tails, and the nested rollout shape.
2026-07-16 20:38:57 -07:00
alexzms 79a5930340 [fix] dreamverse: drop steering hint line that overlapped the hero title on mobile
The two-line 'Drive each scene yourself...' hint grew the vertically-centered
landing content, pushing the preset cards up into the absolutely-positioned
hero title on short mobile viewports. Removing it restores the clean spacing.
2026-07-15 23:20:23 +00:00
alexzms 3a0ce7c20d [misc] dreamverse: bump steering generating progress bar to 10.5s 2026-07-15 23:20:23 +00:00
alexzms 0b006a9e46 [feat] dreamverse: robust steering 'Generating next scene' overlay
Drive the generating indicator off whether the new segment has actually
landed (buffered timeline grows past the boundary captured at generation
start) instead of bare playback events, so scrubbing back and replaying to
the end keeps it visible and it clears promptly once frames arrive. Gate the
overlay on playbackReachedEnd like 'Segment complete', and bump the progress
bar ETA to 8s.
2026-07-15 23:20:23 +00:00
alexzms fb9acac514 [feat] dreamverse: steering-only UI — drop auto-rollout mode selector
Remove the pre-session Auto rollout / Steering segmented control and make
manual continuation (steering) the sole mode: default it on in the session
store and keep it on across project/lobby resets. Leaves a short steering
hint in its place.
2026-07-15 23:20:22 +00:00
kevin314 2d78957219 6+2 2026-07-15 23:20:22 +00:00
kevin314 63991d2017 Test 2026-07-15 23:20:22 +00:00
alexzms b6eafbea50 [fix] dreamverse: recover steering after a blocked or failed prompt
When a steering prompt is refused by the enhancer (e.g. content-policy) or enhancement
fails for all providers, the backend enqueues nothing and won't re-emit prompt_sources_
blocked, leaving the UI stuck on the generating overlay. On prompt/fallback_used and
session/error in steering, return to the 'describe the next scene' state, drop the failed
scene from the history (steeringFailed), and surface a retry notice; clear the notice on
the next submit.
2026-07-15 23:20:22 +00:00
alexzms 7b84041642 [feat] dreamverse: steering scene history of per-segment user prompts
Steering mode now shows an elegant list of each scene's prompt above the player. The
text comes from the user's own words, captured stably at submit time as rawText (the
backend later overwrites text/source with the enhanced prompt, so those are never read);
a preset's opening scene with no user prompt falls back to promptHistory. The redundant
ChatBar 'Segment complete' banner is dropped (the video overlay already says it, and it
was squeezing the list), and a ResizeObserver re-pins the list to the latest scene when
the area resizes.
2026-07-15 23:20:22 +00:00
kevin314 48801c29c7 Add image input UI 2026-07-15 23:20:22 +00:00
alexzms cb8be0d3f3 [feat] dreamverse: main-UI steering mode switch + first-segment-only seeding
Add a user-facing Auto rollout / Steering segmented control to the main composer (not just
devtools), wired to manualContinuationMode and authoritative at session start. In steering
mode a preset seeds only its first segment and auto/loop are forced off, so the backend waits
for the user to describe each subsequent scene by hand.
2026-07-15 23:20:22 +00:00
alexzms 1bc7ff79e6 [ui] dreamverse: logo links home, replace Join Waitlist with Blog
The FastVideo logo now navigates to the app home (most intuitive), and the
Join-Waitlist buttons (header desktop/mobile + session-ended card) become a Blog
link pointing at the Dreamverse blog.
2026-07-15 23:20:21 +00:00
alexzms 76fbe472ae [feat] dreamverse: allow video download at any time during playback
handleDownloadVideo already remuxes live (including in-progress) segments, but the button
was gated on a finalized clip blob. Surface it as soon as playback starts (avPlaybackStarted)
so the user can grab the in-progress video at any point — important for unbounded steering
sessions.
2026-07-15 23:20:21 +00:00
alexzms eba74c43ef [feat] dreamverse: graceful segment-complete & generating overlays in steering playback
When a segment finishes in steering mode the player no longer spins. Instead it shows a
soft 'Segment complete' prompt over the frozen last frame (gated on the playhead actually
reaching the buffered end, and hidden again when the user scrubs back). After the user
submits the next scene, a ~4.5s progress bar covers the generation latency so the wait has
a visible ETA.
2026-07-15 23:20:21 +00:00
alexzms 7744a74c13 [feat] dreamverse: steering-mode toggle in devtools composer
Add a 'Steering mode' checkbox to the devtools composer and thread the
manualContinuationEnabled / onManualContinuationToggle props through DevtoolsShell.
2026-07-15 23:20:21 +00:00
alexzms d746b8b956 [feat] dreamverse: unlimited segments in steering mode
Steering (manual continuation) lets the user drive the rollout segment-by-segment
indefinitely. Treat it like single-clip mode for the generation cap:
_resolve_generation_segment_cap returns 0 (unlimited) and the cap-reached guard is
skipped when manual_continuation_mode is on.
2026-07-15 23:20:21 +00:00
alexzms 65f12605d8 [bugfix] ltx2: apply video_position_offset_sec RoPE offset
The DiT forward swallowed video_position_offset_sec via **kwargs, so multi-segment
rollouts never advanced the temporal RoPE phase between segments, causing ~1s audio/
video desync at each seam. Add the offset to the temporal position coords (mirrors
hao-ai-lab/FastVideo#1422).
2026-07-15 23:20:21 +00:00
kevin314 1e1ac08cd0 Add manual continuation 2026-07-15 23:20:21 +00:00
Mac LeeandSolitaryThinker a253856147 [ci] Add exact identity performance statuses (#1560)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-15 14:59:09 -07:00
Satyam Srivastavaandgemini-code-assist[bot] 6e25d94ebc [ci]: enable scheduled perf runs to update rolling baseline (#1599)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-07-14 16:55:12 -07:00
Mac Lee cae8fa18dc [bugfix]: propagate Qwen2.5-VL visual dtype (#1580) 2026-07-13 18:37:11 -07:00
William Lin 821e5a0832 [bugfix]: fix FlashAttention resolver tests after tuple return (#1597) 2026-07-13 16:29:03 -07:00
William Lin c1abc42782 [bugfix]: allow unrestricted head sizes in SDPA (#1596) 2026-07-13 16:04:51 -07:00
William Lin ef15ea2391 [bugfix]: keep LTX2 rms_norm outputs bf16 under torch 2.12 autocast (#1587) 2026-07-13 16:04:21 -07:00
Mac Lee 1ea2517e22 [ci]: extend LoRA training CI timeout (#1589) 2026-07-13 02:47:28 -07:00
MookandSolitaryThinker 0c63528c59 [perf] Cache RoPE position-embedding tables across denoising steps (#1442)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-13 02:34:51 -07:00
b063f8ca41 [feat] Fix FLUX.1-dev port: native RoPE, parity tests, SSIM reference (#1321)
Co-authored-by: Ishan Vaish <ivaish@ucsd.edu>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-13 01:49:42 -07:00
William Lin e7fff0173a [bugfix]: benchmark_weight_loading_comparison.py — iterate safe_open via .keys() (#1378) 2026-07-12 22:50:45 -07:00
Shreejith SGandH1yori233 d82abc271e [feat] Add GLM-Image inference support (#1030)
Co-authored-by: H1yori233 <k1kong@ucsd.edu>
2026-07-12 22:42:42 -07:00
Guian FangandSolitaryThinker 970409962f [feat] Add AnyFlow any-step video distillation (pretrain + on-policy) (#1371)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-12 02:32:08 +00:00
Raghav K 055586703d [perf]: register a real backward for FA2 default + masked/varlen custom ops (training-under-compile) (#1388) 2026-07-12 00:45:27 +00:00
Mac Lee 5d89f86675 [ci] Stop forcing FA4 in model-load lanes (#1561) 2026-07-11 14:10:36 -07:00
Satyam Srivastava 19a51a1fe6 [ci] Trigger performance benchmarks for performance code changes (#1583) 2026-07-10 20:21:33 -07:00
William Lin d3232cea5a [ci]: gate the full-suite trigger on pre-commit and docs build (#1572) 2026-07-11 02:56:49 +00:00
Raghav KandSolitaryThinker 0c90c8c24d [bugfix] nvfp4: cast fp32 inputs to bf16 instead of asserting (#1488)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-10 21:39:17 +00:00
Mingjia HuoandClaude Fable 5 4c08ffce49 [feat] World model training using third person games (#1443)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-07-10 05:09:07 +00:00
Atharv Ramesh af4a77553c [ci]: add SSIM reference bootstrap flow (#1522) (#1547) 2026-07-10 01:49:38 +00:00
alexzmsandSolitaryThinker c096fda1eb [docs] Add LTX-2.3 distilled inference run configs (t2v/i2v × 5+2/8+3 × resolutions) (#1568)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-09 18:53:12 +00:00
William Lin 8f47e85be0 [bugfix]: retry remote image downloads in load_image (#1570) 2026-07-09 07:33:24 -07:00
William Lin 90d3bd19eb [infra] Deliver per-job Buildkite env to Modal CI at runtime, not as image layers (#1569) 2026-07-09 06:33:37 -07:00
Mac LeeandSolitaryThinker afb4f7d3c5 [ci]: emit v2 performance result schema (#1551)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-09 06:11:15 +00:00
02e1143f22 [feat] Add Kandinsky-5 T2V/I2V pipeline support (#1471)
Co-authored-by: Aryan Kumar <aryan5v@users.noreply.github.com>
Co-authored-by: leffff <levnovitskiy@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 14:43:30 -07:00
KaredandSolitaryThinker e2f4d1a7b5 [feat]: add SwanLab tracker (#1461)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 08:35:38 +00:00
Kaiqin KongandSolitaryThinker f037351146 [feat] Add Clean-history Teacher Forcing and Causal Consistency Distillation (#1505)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 08:13:11 +00:00
Mac LeeandSolitaryThinker 1ee11e08dc [ci]: add performance fingerprint cohorts (#1546)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 07:24:00 +00:00
595f0ea60e [feat] Add DreamX-World 5B Cam and AR pipelines (#1538)
Co-authored-by: Suckl <Suckl@users.noreply.github.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 06:16:18 +00:00
William Lin d921832cd2 [misc]: reserve CPU/memory for timing-sensitive Modal CI lanes (#1566) 2026-07-06 22:13:53 -07:00
William Lin 629697629a [bugfix]: free CUDA memory between train-framework model tests (LongCat OOM on L40S) (#1565) 2026-07-06 21:27:30 -07:00
William Lin a25313beec [ci]: wire fastvideo/tests/ops/ into the unit-test lane (#1559) 2026-07-06 12:07:26 -07:00
William Lin dbde64385b [bugfix]: bump FA4 pin to the CuTe DSL 4.6 compatible rev (#1564) 2026-07-06 12:06:59 -07:00
William Lin 9d909f5f04 [test]: remove dead and duplicate tests (-489 lines) (#1556) 2026-07-05 15:53:40 -07:00
William Lin 76b0550c15 [ci]: run pre-commit on fork PRs without manual approval (#1555) 2026-07-05 14:18:16 -07:00
William Lin 384c1e9493 [misc]: update reseed-performance-baseline skill for the hf_store move (#1545 follow-up) (#1553) 2026-07-05 14:16:55 -07:00
William Lin b1dbcc93f6 [misc]: reformat fastvideo/performance to the repo yapf config (#1554) 2026-07-05 14:16:20 -07:00
William Lin b93833772e [ci]: guard against test directories no CI lane collects (#1552) 2026-07-05 14:07:38 -07:00
Mac Lee 30b523edd6 [ci] Normalize performance stage component metrics (#1475) (#1550) 2026-07-05 14:05:26 -07:00
Mac LeeandSolitaryThinker 6aab7f3832 [ci] cover Hunyuan 1.5 chat-list text preprocessing (#1518)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 12:05:25 -07:00
Mac Lee 9cd53fe5f8 [ci] Add metric-specific performance thresholds (#1545) 2026-07-05 12:04:33 -07:00
Mac Lee 6a32cf3a5e [ci]: expose LoRA extraction slash command (#1542) 2026-07-05 06:45:31 -07:00
William Lin 98be9b3da2 [bugfix]: address the three remaining #1447 review findings (#1549) 2026-07-05 06:43:27 -07:00
Mac Lee c53e85b767 [ci] Add v2 performance benchmark config identity fields (#1544) 2026-07-05 06:18:15 -07:00
William Lin 40a8bd2d3b [bugfix]: skip ThunderKittens kernels on aarch64 and document the kernel build matrix (#1548) 2026-07-05 06:13:11 -07:00
Mac LeeandSolitaryThinker 31aa115611 [bugfix]: preserve FSDP hooks for RMSNorm qk norms (#1513)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 04:15:49 +00:00
zainnhandSolitaryThinker 98ac10a528 [infra] Auto-rebuild CUDA images when docker/Dockerfile changes on main (#1526)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-04 16:25:46 -07:00
William Lin a5a6d171e5 [attn] Make FA4 explicit opt-in via FASTVIDEO_FA4 and delete the runtime fallback machinery (#1540) 2026-07-04 14:51:00 -07:00
327 changed files with 34656 additions and 1704 deletions
@@ -1,20 +1,23 @@
---
name: reseed-performance-baseline
description: Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift using one or more reviewed normalized performance JSONs. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline in FastVideo/performance-tracking must be advanced from a consistent batch of reviewed source results. The workflow backs up existing history under /tmp, validates all source JSONs for the same (model_id, gpu_type), rejects internally inconsistent source batches, uploads one success=true reseed record per accepted source JSON, and offers to clean local temp state after a successful upload.
description: Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, environment-caused benchmark shift, or reviewed v2 calibration using one or more reviewed normalized performance JSONs. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline in FastVideo/performance-tracking must be advanced, or when a new v2 exact comparable identity needs its first approved baseline. The workflow backs up existing history under /tmp, validates all source JSONs for the same legacy (model_id, gpu_type) target or the same v2 exact identity, rejects internally inconsistent source batches, uploads one success=true baseline record per accepted source JSON, and offers to clean local temp state after a successful upload.
---
# Re-seed Performance Baseline
## Purpose
Replace or advance the rolling performance baseline for a single
`(model_id, gpu_type)` pair in the HF dataset
`FastVideo/performance-tracking`.
Replace or advance the rolling performance baseline in the HF dataset
`FastVideo/performance-tracking`. Legacy targets are scoped by
`(model_id, gpu_type)`. V2 targets are scoped by exact comparable identity:
`workload_id`, `variant_id`, `benchmark_version`, `hardware_profile_id`,
`software_profile_id`, and `recipe_fingerprint`.
Performance comparison uses the median of up to the last 5 successful records
for the same model and GPU. Failed records are useful audit history, but they
do not move the future baseline because `compare_baseline.py` loads records
with `successful_only=True`.
Performance comparison uses the median of up to the last 5 successful,
baseline-eligible records for the same target. Failed or calibration-only
records are useful audit history, but they do not move the future baseline
because `compare_baseline.py` loads records with `successful_only=True` and
`baseline_eligible_only=True`.
This skill now reseeds from a reviewed batch of one or more source performance
JSONs. It uploads one new `success=true` record per accepted source JSON; it
@@ -22,11 +25,13 @@ does not blindly replicate one measurement into 3 or 5 records. The effective
reseed size is therefore dynamic and equals the number of provided, validated,
internally consistent source JSONs.
If the operator provides fewer than 3 records, call out that the last-5 rolling
median may not move immediately. If the operator provides 3 consistent shifted
records, the rolling median usually moves immediately. If the operator provides
5 consistent shifted records, the last-5 window is effectively reset to the new
runtime profile.
For baseline shifts with existing history, if the operator provides fewer than
3 records, call out that the last-5 rolling median may not move immediately. If
the operator provides 3 consistent shifted records, the rolling median usually
moves immediately. If the operator provides 5 consistent shifted records, the
last-5 window is effectively reset to the new runtime profile. For the first
approved v2 baseline of a new exact identity, one reviewed calibration seed is
enough for the next comparable run to leave `CALIBRATION_NEEDED`.
These records are intentional operator-approved baseline resets, not ordinary
independent main-branch persistence. Mark them clearly with provenance fields
@@ -66,10 +71,10 @@ approval, then upload reviewed accepted baseline records.
| Parameter | Required | Description |
|-----------|----------|-------------|
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
| `model_id` | Legacy required; v2 inferred | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. For v2 records, use the `model_id` from each source artifact only as the upload directory; comparison is by exact identity. |
| `gpu_type` | Legacy required; v2 inferred | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. V2 hardware matching uses `hardware_profile_id`; preserve `gpu_type` as display metadata. |
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
Hardcoded defaults:
@@ -78,10 +83,14 @@ Hardcoded defaults:
supported by the code, but use the default unless the user explicitly asks).
- Local sync root: `/tmp/perf-tracking` (`PERFORMANCE_TRACKING_ROOT` override
is supported).
- Prepared-record staging root: `/tmp/performance_reseed_prepared`
(`PERFORMANCE_RESEED_STAGING_ROOT` override is supported). Keep it separate
and non-nested from the sync root.
- Backup root: `/tmp/performance_reseed_backup`.
- Download scratch root for source artifact URLs: `/tmp/performance_reseed_source`.
- Baseline window: last 5 `success=true` records for the same
`(model_id, gpu_type)`.
- Baseline window: last 5 `success=true`, `baseline_eligible=true` records
for the same legacy `(model_id, gpu_type)` target or the same v2 exact
comparable identity.
- Reseed count: dynamic. Upload exactly one accepted seed record per validated
source JSON.
@@ -115,12 +124,24 @@ with open(source_result, encoding="utf-8") as f:
record = json.load(f)
```
Stop if any normalized record's `model_id` or `gpu_type` does not match the
requested `model_id` and `gpu_type`.
Classify the source batch before continuing:
- **Legacy source records** have no v2 exact identity fields. Stop if any
normalized record's `model_id` or `gpu_type` does not match the requested
`model_id` and `gpu_type`.
- **V2 source records** have exact identity fields. Stop unless every source
record has all six comparable identity fields and they are identical across
the batch: `workload_id`, `variant_id`, `benchmark_version`,
`hardware_profile_id`, `software_profile_id`, and `recipe_fingerprint`.
Do not fall back to legacy `(model_id, gpu_type)` matching for v2 records.
The source records may have `success: false` when they came from failed
rolling baseline comparisons. That is expected; only the reviewed reseed
records become new `success: true` baseline records after explicit approval.
For a first v2 baseline seed, the source records must instead be successful
scheduled-main full-suite `CALIBRATION_NEEDED` normalized artifacts. Reject PR,
local, direct-run, non-main-branch, or non-full-suite calibration artifacts as
seed sources.
Sort validated source records by their original `timestamp` ascending before
preparing the seed records. If a source timestamp is missing or unparsable,
@@ -148,8 +169,7 @@ For each metric with at least two non-null source values:
4. Stop if any source record regresses against the batch median by more than
`max_intra_batch_regression`.
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
otherwise `0.05`. Print a table with per-source values, batch median, and
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
worst intra-batch regression.
This check prevents uploading a mixed batch where one JSON is materially
@@ -183,7 +203,7 @@ present, that run is not a valid source for baseline reseeding.
### 2. Sync and back up existing HF records under /tmp
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
records, while this step must only fetch and back up existing history.
@@ -192,16 +212,16 @@ The sync command pattern is:
```bash
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
```
Then back up only the sanitized model directory under `/tmp`:
For legacy records, back up the sanitized model directory under `/tmp`:
```bash
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
from hf_store import sanitize
MODEL_SAFE=$(python - <<'PY'
from fastvideo.performance.hf_store import sanitize
print(sanitize("<model_id>"))
PY
)
@@ -210,6 +230,16 @@ mkdir -p "$BACKUP_DIR"
cp -R "${PERFORMANCE_TRACKING_ROOT}/${MODEL_SAFE}" "$BACKUP_DIR/" 2>/dev/null || true
```
For v2 records, back up the full local tracking root after sync. Exact identity
lookup scans across model directories, so a benchmark rename may have relevant
history outside the current source artifact's `model_id` directory:
```bash
BACKUP_DIR="/tmp/performance_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_v2_exact_identity"
mkdir -p "$BACKUP_DIR"
cp -R "${PERFORMANCE_TRACKING_ROOT}" "$BACKUP_DIR/tracking-root"
```
Write provenance next to the backup:
```bash
@@ -232,10 +262,12 @@ first baseline seed. Continue, but report that baseline history was empty.
### 3. Compute old baseline and candidate shift
Load the last 5 successful records for the target:
Load the last 5 successful baseline records for the target.
For legacy targets:
```python
from hf_store import load_records_for_model
from fastvideo.performance.hf_store import load_records_for_model
records = load_records_for_model(
"/tmp/perf-tracking",
@@ -243,6 +275,28 @@ records = load_records_for_model(
"<gpu_type>",
last_n=5,
successful_only=True,
baseline_eligible_only=True,
)
```
For v2 exact-identity targets:
```python
from fastvideo.performance.hf_store import load_records_for_identity
records = load_records_for_identity(
"/tmp/perf-tracking",
{
"workload_id": "<workload_id>",
"variant_id": "<variant_id>",
"benchmark_version": "<benchmark_version>",
"hardware_profile_id": "<hardware_profile_id>",
"software_profile_id": "<software_profile_id>",
"recipe_fingerprint": "<recipe_fingerprint>",
},
last_n=5,
successful_only=True,
baseline_eligible_only=True,
)
```
@@ -258,7 +312,8 @@ medians after appending the proposed seed records, and source batch spread for:
Also print how many successful old records exist. Make clear:
- 1 seed record usually does not move a last-5 median by itself.
- 1 seed record usually does not move an existing last-5 median by itself, but
it is enough to establish the first v2 baseline for a new exact identity.
- 3 consistent seed records usually move the last-5 median immediately.
- 5 consistent seed records effectively reset the last-5 window.
- The records are intentional approved baseline resets and must be labeled
@@ -268,10 +323,10 @@ Also print how many successful old records exist. Make clear:
Require an explicit confirmation phrase before preparing the upload:
> About to RE-SEED performance baseline for `<model_id>` on `<gpu_type>`.
> About to RE-SEED performance baseline for `<target description>`.
> This will upload `<N>` new `success=true` records to
> `FastVideo/performance-tracking/<sanitize(model_id)>/`, one per accepted
> source JSON.
> `FastVideo/performance-tracking/<sanitize(model_id)>/` or the source
> artifact's v2 model directory, one per accepted source JSON.
>
> Reason: `<intent_rationale>`
> Source results: `<source_results>`
@@ -289,25 +344,63 @@ Do not continue unless the user types exactly `confirm performance reseed`.
### 5. Create the accepted seed records
Create one seed record from each normalized source result. Do not copy the
Create one seed record from each normalized source result.
For first v2 baseline seeds, use the scoped utility. It validates exact
identity, requires successful scheduled-main full-suite `CALIBRATION_NEEDED`
source artifacts, preserves the normalized v2 identity and metadata fields,
and writes seed records with `success=true`, `baseline_eligible=true`, and
`comparison_status=PASS`:
```bash
python fastvideo/tests/performance/seed_baseline.py \
--source-result <normalized_perf_1.json> \
--source-result <normalized_perf_2.json> \
--intent-rationale "<intent_rationale>" \
--max-intra-batch-regression 0.05 \
--tracking-root "${PERFORMANCE_TRACKING_ROOT}" \
--staging-root "${PERFORMANCE_RESEED_STAGING_ROOT:-/tmp/performance_reseed_prepared}"
```
The utility is prepare-only and intentionally has no upload option. Upload the
scoped records only after the separate confirmation in step 6.
The utility validates against an isolated fresh HF snapshot and leaves
`PERFORMANCE_TRACKING_ROOT` untouched; that argument only proves the staging
root is separate from the operator's tracking mirror. Before writing, it stops
if the exact identity already has a successful baseline-eligible record or if
the workload/variant/version already trusts another recipe. It atomically
reserves the exact identity and writes a digest-protected upload manifest bound
to the current HF endpoint, repository id, and repository type. Keep the
prepared records, manifest, source files, and reservation unchanged until the
operation is uploaded or explicitly cleaned up.
If the prepared seed records look correct, upload only those scoped records in
step 7. Do not rerun the utility with a different source list after approval.
For legacy reseeds or accepted v2 baseline shifts from regression artifacts,
create one seed record from each normalized source result. Do not copy the
source JSON wholesale.
Infer the baseline field allowlist from all existing HF records for the target
`(model_id, gpu_type)` after syncing, including both `success=true` and
`success=false` records. Use the union of non-provenance keys present in those
target records, preserving only fields that also exist in the normalized
source record or are explicitly set by the reseed workflow. Always include
`model_id`, `timestamp`, and `success` because the upload path and baseline
loader depend on them. Always set `timestamp` to a fresh reseed timestamp and
`success` to `true`. Do not include unrelated source-only fields that are
absent from existing HF records.
after syncing, including both `success=true` and `success=false` records. For
legacy targets the target is `(model_id, gpu_type)`. For v2 baseline-shift
reseeds the target is the exact comparable identity. Use the union of
non-provenance keys present in those target records, preserving only fields
that also exist in the normalized source record or are explicitly set by the
reseed workflow. Always include `model_id`, `timestamp`, `success`,
`baseline_eligible`, and `comparison_status` because the upload path and
baseline loader depend on them. Always set `timestamp` to a fresh reseed
timestamp, `success` to `true`, `baseline_eligible` to `true`, and
`comparison_status` to `PASS`. Do not include unrelated source-only fields
that are absent from existing HF records.
Exclude existing provenance or operator metadata from the inferred baseline
field allowlist. At minimum, exclude keys prefixed with `baseline_reseed` and
any fields known to be local-only audit metadata.
If there are no previous HF records for the target model/GPU, fall back to this
default baseline field list:
If there are no previous HF records for the target, fall back to this default
baseline field list:
- `model_id`
- `timestamp`
@@ -320,6 +413,22 @@ default baseline field list:
- `dit_time_s`
- `vae_decode_time_s`
- `success`
- `baseline_eligible`
- `comparison_status`
For v2 baseline-shift reseeds with no previous HF records for the exact
identity, also preserve:
- `workload_id`
- `variant_id`
- `benchmark_version`
- `recipe_fingerprint`
- `hardware_profile_id`
- `software_profile_id`
- `recipe`
- `hardware_profile`
- `software_profile`
- `software_comparison_profile`
Do not upload extra fields from the source artifact.
@@ -335,6 +444,22 @@ Optional provenance fields are allowed and useful:
- `baseline_reseed_operator`
- `baseline_reseed_max_intra_batch_regression`
The v2 calibration seed utility writes analogous first-seed provenance:
- `baseline_seed: true`
- `baseline_seed_reason`
- `baseline_seed_source_result`
- `baseline_seed_source_status`
- `baseline_seed_source_timestamp`
- `baseline_seed_source_success`
- `baseline_seed_source_run_source`
- `baseline_seed_source_branch`
- `baseline_seed_source_test_scope`
- `baseline_seed_source_pr_number`
- `baseline_seed_batch_size`
- `baseline_seed_batch_index`
- `baseline_seed_operator`
Use a fresh reseed timestamp for each seed record, not the original source
result timestamp. This is required because
`load_records_for_model(..., last_n=5)` keeps the last records after loading
@@ -357,7 +482,8 @@ Prefer uploading new accepted seed records so failed history remains visible.
Print:
- Backup directory path under `/tmp`.
- Prepared local record paths under `PERFORMANCE_TRACKING_ROOT`.
- Prepared local record paths under `PERFORMANCE_RESEED_STAGING_ROOT`.
- Prepared upload-manifest path under the identity reservation.
- HF paths that will receive the new records.
- Old rolling medians.
- Source batch medians, source batch spread, reseed count, and candidate
@@ -369,22 +495,36 @@ prepared records plus backup on disk.
### 7. Upload only the scoped records
Use the shared storage helper so the path and repo type match CI:
For a first v2 calibration seed, use the manifest uploader after the user
replies exactly `upload`:
```python
from hf_store import upload_record
upload_record("<local_record_path>", record, strict=True)
```bash
python -c 'from fastvideo.tests.performance.seed_baseline import upload_prepared_seed_manifest; print(upload_prepared_seed_manifest("<prepared_manifest>"))'
```
Run it once per prepared record. Each upload goes to:
The uploader verifies the source and prepared-record digests, pins and scans
the current HF revision, rechecks exact-identity and recipe-cohort conflicts,
and writes the entire batch in one commit whose `parent_commit` must still be
current. A concurrent Hub update makes the commit fail. Do not retry
automatically: preserve staging, refresh/review remote state, and request a new
explicit `upload` after the conflict is understood. Each record goes to:
```text
FastVideo/performance-tracking/<sanitize(model_id)>/<record_filename>.json
```
Never bulk upload the whole tracking root. Never modify another model's
directory in the same operation.
Never call `upload_record()` once per first-seed record: that can partially
land the batch and has no compare-and-swap guard.
For a legacy reseed or an accepted v2 baseline shift, the first-seed manifest
validator does not apply because an eligible baseline already exists. Upload
only the individually reviewed records prepared in step 5 with the shared
`upload_record(local_path, record, strict=True)` helper. Stop on the first
failure and report exactly which records reached HF; do not silently rerun or
replicate the remainder.
Never bulk upload the tracking or staging root, and never modify another
model's directory in the same operation.
### 8. Report outcome and offer cleanup
@@ -406,9 +546,14 @@ distinguish an accepted baseline shift from a hidden regression.
After the upload is verified, ask whether the user wants to clear temporary
local state. Explain what each directory is for:
- `PERFORMANCE_TRACKING_ROOT`, usually `/tmp/perf-tracking`: local synced
mirror of `FastVideo/performance-tracking` plus the prepared local seed
records used for scoped upload.
- `PERFORMANCE_TRACKING_ROOT`, usually `/tmp/perf-tracking`: read-only local
synced mirror used for operator review and reporting. First-v2 preparation
independently proves remote state from a fresh temporary HF snapshot.
- `PERFORMANCE_RESEED_STAGING_ROOT`, usually
`/tmp/performance_reseed_prepared`: prepared local seed records used for the
scoped upload, plus the identity reservation and digest manifest. Keeping
this separate prevents aborted preparations from appearing in later
baseline reads.
- `/tmp/performance_reseed_backup/<...>`: local backup of the target model's
pre-reseed HF history plus `PROVENANCE.txt`, kept so a bad reseed can be
audited or corrected.
@@ -418,14 +563,18 @@ local state. Explain what each directory is for:
Ask:
> Reseed succeeded. Do you want me to delete the local temp tracking mirror,
> source downloads, and reseed backup under `/tmp`? These files are local
> safety/audit artifacts only; HF already has the uploaded records.
> this reseed's prepared staging records, source downloads, and reseed backup
> under `/tmp`? These files are local safety/audit artifacts only; HF already
> has the uploaded records.
>
> Reply `cleanup reseed temp` to delete them, anything else to keep them.
Do not delete anything unless the user replies exactly
`cleanup reseed temp`. If cleanup is requested, remove only the specific
directories created for this reseed. Never remove unrelated `/tmp` contents.
directories and prepared record paths created for this reseed. Do not remove
the shared staging root when it contains other records. Remove this operation's
identity reservation only with its prepared records and manifest, and never
remove unrelated `/tmp` contents.
## Failure modes and handling
@@ -437,19 +586,34 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
against the source batch median by more than `max_intra_batch_regression`.
Ask for cleaner sources or a reviewed explanation before continuing.
- **Too few source records to move the median.** Continue only after making
clear that one or two records may not immediately move the last-5 median.
clear that one or two records may not immediately move an existing last-5
median. This warning does not block a first v2 calibration seed for an exact
identity with no eligible baseline yet.
- **The source results are noisy or suspicious.** Stop. Reseeding amplifies
those measurements into the baseline, so they must be reviewed first.
- **HF sync fails.** Stop for destructive reseeds. A stale or empty sync can
make the old baseline look missing.
- **The exact v2 identity already has an eligible baseline.** Stop. The
`CALIBRATION_NEEDED` artifact is stale; use the reviewed baseline-shift path
instead of the first-seed utility.
- **The workload/variant/version trusts another recipe.** Stop. The source is
stale relative to the current recipe cohort and must not bypass
`RECIPE_MISMATCH` by creating a second trusted recipe.
- **The staging root already has a prepared seed for the exact identity.**
Stop and reuse, upload, or explicitly clean that preparation. Do not prepare
another copy of the same measurement.
- **The conditional Hub commit loses its parent race.** Stop without retrying.
Keep the preparation, refresh and review the new remote state, then request
a new explicit `upload` only if the seed is still valid.
- **Candidate still violates fixed thresholds.** Report that this skill only
handles the rolling HF baseline; update benchmark JSON thresholds in code
review if maintainers accept the new absolute limit.
- **The user aborts at either confirmation.** Leave the backup and prepared
records on disk. Nothing should be uploaded.
- **The user declines cleanup.** Keep `/tmp/perf-tracking`, the source
download directory if any, and `/tmp/performance_reseed_backup/<...>` in
place for audit/debugging.
- **The user declines cleanup.** Keep `/tmp/perf-tracking`, the prepared seed
records under `/tmp/performance_reseed_prepared`, the source download
directory if any, and `/tmp/performance_reseed_backup/<...>` in place for
audit/debugging.
- **A bad seed was uploaded.** Use the backup and HF history to identify the
uploaded file, then remove or supersede it with an explicitly reviewed
corrective record. Do not silently rewrite unrelated history.
@@ -460,8 +624,9 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
intentional baseline replacement.
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
median comparison, and persistence rules.
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
`sanitize()`, and `upload_record()`.
- `fastvideo/performance/hf_store.py` — HF sync and record loading helpers.
- `fastvideo/tests/performance/seed_baseline.py` — first-seed preparation,
staging reservation, manifest validation, and conditional batch upload.
- `fastvideo/tests/performance/test_inference_performance.py` — source result
JSON schema.
- `.buildkite/performance-benchmarks/tests/*.json` — fixed absolute benchmark
@@ -474,3 +639,4 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
| 2026-05-03 | Initial version. Sister workflow to `reseed-ssim-references`, scoped to one performance `(model_id, gpu_type)` baseline seed with backup, confirmation, provenance, and `success=true` upload. |
| 2026-05-03 | Previous policy: replicate one approved shifted source result into 3 success records by default, or 5 only when explicitly requested. Add provenance marker for replicated-source reseeds. Superseded by the 2026-05-08 dynamic multi-source policy. |
| 2026-05-08 | Replace fixed 3/5 replication with dynamic multi-source reseeding: upload one seed record per reviewed source JSON, validate intra-batch consistency, move backup/source scratch under `/tmp`, and ask whether to clean temp state after successful upload. |
| 2026-07-13 | Keep first-v2-seed preparation outside the canonical mirror, reserve staging identities atomically, reject stale or replayed calibration seeds, and upload reviewed manifests with a single parent-guarded Hub commit. |
@@ -1,5 +1,9 @@
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 3,
"description": "Wan2.1 T2V 1.3B inference performance",
"model": {
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
+28 -1
View File
@@ -114,6 +114,17 @@ steps:
limit: 2
agents:
queue: "default"
- label: ":test_tube: LoRA Extraction Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Training Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
@@ -371,6 +382,21 @@ steps:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "scripts/lora_extraction/**"
- "fastvideo/tests/lora_extraction/**"
- "fastvideo/models/loader/**"
- "fastvideo/training/training_utils.py"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
@@ -410,7 +436,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Training Tests"
env:
- TEST_TYPE=training_lora
@@ -455,6 +481,7 @@ steps:
- "fastvideo/layers/**"
- "fastvideo/worker/**"
- "fastvideo/entrypoints/**"
- "fastvideo/performance/**"
- "fastvideo/tests/performance/**"
- ".buildkite/performance-benchmarks/**"
- "pyproject.toml"
+24 -2
View File
@@ -76,10 +76,27 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
EFFECTIVE_PR=$PR_NUMBER
fi
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} BUILDKITE_SOURCE=${BUILDKITE_SOURCE:-} TEST_SCOPE=${TEST_SCOPE:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
POST_RUN_HOOK=""
is_truthy() {
case "${1:-}" in
1|true|TRUE|yes|YES|on|ON) return 0 ;;
*) return 1 ;;
esac
}
ssim_bootstrap_args() {
local title="${PR_TITLE:-}"
local message="${BUILDKITE_MESSAGE:-}"
if is_truthy "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-}" \
|| [[ "$title" == *"[new-model]"* ]] \
|| [[ "$message" == *"[new-model]"* ]]; then
printf ' --bootstrap-mode'
fi
}
upload_performance_artifacts() {
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
LOCAL_DIR="downloaded_reports"
@@ -172,7 +189,12 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
if [ -n "$SSIM_BOOTSTRAP_ARGS" ]; then
log "SSIM bootstrap mode enabled for new-model reference draft generation"
fi
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run "
MODAL_COMMAND+="$MODAL_SSIM_TEST_FILE::run_ssim_tests$SSIM_BOOTSTRAP_ARGS"
;;
"training")
log "Running training tests..."
+133
View File
@@ -0,0 +1,133 @@
#!/usr/bin/env bash
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
#
# Polls the workflow runs for the PR head commit and only exits 0 once the
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
# already doomed.
#
# Semantics:
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
# no full suite; the next push re-arms via the 'synchronize' trigger)
# - watched run cancelled -> still pending: the docs
# workflow's repo-global 'pages' concurrency group cancels runs superseded
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
# - watched runs pending -> poll until done
# - docs run absent -> not applicable after a
# short grace period ('Deploy Documentation' is path-filtered on PRs)
# - pre-commit run absent -> keep polling: pre-commit
# is never path-filtered, so its absence is always anomalous
# - 'ready' label removed while waiting -> exit 1 (fail CLOSED:
# un-labeling is a deliberate maintainer action)
# - GitHub API unreachable or timeout -> exit 0 (fail OPEN,
# loud warning: never brick CI on a GitHub outage)
#
# Required env: PR_SHA (PR head commit), PR_NUMBER, GITHUB_REPOSITORY, GH_TOKEN.
set -euo pipefail
: "${PR_SHA:?PR_SHA (PR head commit) is required}"
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
# Workflow-level `name:` values that must be green before the full suite
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
# legitimately never exist; pre-commit always runs, so it must appear.
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
WATCHED_REGEX='^(pre-commit|Deploy Documentation)$'
POLL_SECS="${POLL_SECS:-20}"
GRACE_SECS="${GRACE_SECS:-60}"
MAX_WAIT_SECS="${MAX_WAIT_SECS:-1500}"
# Bound each API call so a hung connection hits the 3-strike fail-open path
# instead of pinning the loop until the job timeout (which would fail closed
# on exactly the GitHub-outage case this script is meant to survive).
if command -v timeout >/dev/null 2>&1; then
gh_api() { timeout 30 gh api "$@"; }
else
gh_api() { gh api "$@"; } # macOS dev boxes; CI always has coreutils timeout
fi
# The workflow checked the label before starting the gate, but the wait can
# last ~25 min: re-check once before any exit 0 and fail closed if 'ready'
# was removed in the meantime. An API error here proceeds (the label was
# present when the gate started; never brick CI on an outage).
recheck_ready_label() {
local pr_json
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
exit 1
fi
else
echo "::warning::Could not re-check the 'ready' label on PR #${PR_NUMBER}; proceeding (it was present when the gate started)."
fi
}
start=$(date +%s)
api_fails=0
missing=""
while true; do
elapsed=$(( $(date +%s) - start ))
if runs_json=$(gh_api "repos/${GITHUB_REPOSITORY}/actions/runs?head_sha=${PR_SHA}&per_page=100" 2>/dev/null) \
&& state=$(jq --arg re "$WATCHED_REGEX" '
[.workflow_runs[]? | select(.name // "" | test($re))]
| group_by(.name) | map(max_by(.id))
| map({name, status, conclusion})' <<<"$runs_json" 2>/dev/null); then
api_fails=0
echo "t+${elapsed}s watched checks: $(jq -c . <<<"$state")"
failed=$(jq -r '[.[] | select(.status == "completed"
and (.conclusion | IN("success", "skipped", "neutral", "cancelled") | not))]
| map(.name) | join(", ")' <<<"$state")
if [ -n "$failed" ]; then
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
"label re-arms on every push), or re-run the failed check and then" \
"re-run this workflow."
exit 1
fi
# 'cancelled' counts as pending: wait for a re-run to reach a real verdict
# (bounded by MAX_WAIT, then the fail-open below).
pending=$(jq '[.[] | select(.status != "completed" or .conclusion == "cancelled")] | length' <<<"$state")
missing=$(jq -r --argjson watched "$WATCHED_NAMES" '($watched - map(.name)) | join(", ")' <<<"$state")
if [ "$pending" -eq 0 ]; then
if [ -z "$missing" ]; then
recheck_ready_label
echo "All watched cheap checks are green — full suite may proceed."
exit 0
fi
case "$missing" in
*pre-commit*)
echo "pre-commit run not found for ${PR_SHA} yet; waiting (pre-commit is never path-filtered, so its absence is anomalous)."
;;
*)
if [ "$elapsed" -ge "$GRACE_SECS" ]; then
recheck_ready_label
echo "::warning::Watched run(s) never appeared for ${PR_SHA}: ${missing} (path-filtered, likely not applicable). Proceeding on the checks that did run."
exit 0
fi
echo "Waiting up to ${GRACE_SECS}s grace for path-filtered run(s) to appear: ${missing}."
;;
esac
fi
else
api_fails=$(( api_fails + 1 ))
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
if [ "$api_fails" -ge 3 ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
exit 0
fi
fi
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
exit 0
fi
sleep "$POLL_SECS"
done
+122
View File
@@ -0,0 +1,122 @@
#!/usr/bin/env bash
# Self-test for gate_full_suite.sh using a mocked `gh`. No network, runs on
# any dev box: bash .github/scripts/test_gate_full_suite.sh
set -u
here=$(cd "$(dirname "$0")" && pwd)
tmp=$(mktemp -d)
trap 'rm -rf "$tmp"' EXIT
# Mock gh. Asserts the exact endpoint (including head_sha) it is called
# with — an endpoint typo in the gate script fails the test rather than
# silently serving canned data. On the runs endpoint it serves
# $MOCK_DIR/response_<call#>.json, sticking on the highest existing file,
# and exits 1 if none exist (simulates a GitHub API outage). On the pulls
# endpoint it serves $MOCK_DIR/pr.json, defaulting to a 'ready'-labeled PR.
cat > "$tmp/gh" <<'EOF'
#!/usr/bin/env bash
if [ "${1:-}" != "api" ]; then
echo "unexpected gh invocation: $*" >> "$MOCK_DIR/endpoint_error"
exit 2
fi
case "${2:-}" in
"repos/o/r/actions/runs?head_sha=deadbeef&per_page=100")
n=$(( $(cat "$MOCK_DIR/count" 2>/dev/null || echo 0) + 1 ))
echo "$n" > "$MOCK_DIR/count"
while [ "$n" -gt 0 ]; do
if [ -f "$MOCK_DIR/response_$n.json" ]; then
cat "$MOCK_DIR/response_$n.json"
exit 0
fi
n=$(( n - 1 ))
done
echo "api outage" >&2
exit 1
;;
"repos/o/r/pulls/42")
if [ -f "$MOCK_DIR/pr.json" ]; then
cat "$MOCK_DIR/pr.json"
else
echo '{"labels": [{"name": "ready"}]}'
fi
;;
*)
echo "unexpected gh endpoint: $2" >> "$MOCK_DIR/endpoint_error"
exit 2
;;
esac
EOF
chmod +x "$tmp/gh"
PC_OK='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "success"}'
PC_BAD='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "failure"}'
PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion": null}'
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
fails=0
want_log="" # optional: expect() also greps out.log for this regex, then resets
pr_json="" # optional: served for the pulls (label re-check) endpoint, then resets
raw_body="" # optional: serve responses verbatim instead of wrapping in workflow_runs
expect() { # <name> <expected-exit> <response json>...
local name=$1 want=$2 dir i=1
shift 2
dir=$(mktemp -d "$tmp/test_XXXXXX")
for body in "$@"; do
if [ -n "$raw_body" ]; then
printf '%s' "$body" > "$dir/response_$i.json"
else
printf '{"workflow_runs": [%s]}' "$body" > "$dir/response_$i.json"
fi
i=$(( i + 1 ))
done
[ -n "$pr_json" ] && printf '%s' "$pr_json" > "$dir/pr.json"
( export PATH="$tmp:$PATH" MOCK_DIR="$dir" PR_SHA=deadbeef PR_NUMBER=42 \
GITHUB_REPOSITORY=o/r POLL_SECS=0 GRACE_SECS=1 MAX_WAIT_SECS=3
bash "$here/gate_full_suite.sh" > "$dir/out.log" 2>&1 )
local rc=$?
if [ "$rc" -ne "$want" ]; then
echo "FAIL: $name (exit $rc, want $want)"
cat "$dir/out.log"
fails=1
elif [ -f "$dir/endpoint_error" ]; then
echo "FAIL: $name (mock gh got an unexpected call)"
cat "$dir/endpoint_error"
fails=1
elif [ -n "$want_log" ] && ! grep -Eq "$want_log" "$dir/out.log"; then
echo "FAIL: $name (log does not match: $want_log)"
cat "$dir/out.log"
fails=1
else
echo "ok: $name"
fi
want_log="" pr_json="" raw_body=""
}
expect "both green -> proceed" 0 "$PC_OK, $DOCS_OK, $OTHER, $NULL_NAME"
expect "docs build failed -> blocked" 1 "$PC_OK, $DOCS_BAD"
expect "pre-commit failed -> blocked" 1 "$PC_BAD"
expect "pending then green -> proceed" 0 "$PC_PENDING" "$PC_OK, $DOCS_OK"
want_log="never appeared.*Deploy Documentation"
expect "docs run absent (path-filtered) -> proceed after grace" 0 "$PC_OK"
expect "API outage -> fail open" 0
want_log="FAILING OPEN"
expect "pending past MAX_WAIT -> fail open" 0 "$PC_PENDING"
want_log="FAILING OPEN"
expect "unrelated runs only -> no grace, fail open at MAX_WAIT" 0 "$OTHER"
expect "cancelled docs then green -> proceed" 0 \
"$PC_OK, $DOCS_CANCELLED" "$PC_OK, $DOCS_OK"
want_log="FAILING OPEN"
expect "cancelled docs forever -> fail open at MAX_WAIT" 0 "$PC_OK, $DOCS_CANCELLED"
want_log="FAILING OPEN"
expect "pre-commit absent -> no grace, fail open at MAX_WAIT" 0 "$DOCS_OK"
expect "duplicate run names -> latest wins" 0 "$PC_BAD, $PC_OK_RERUN, $DOCS_OK"
raw_body=1
expect "garbage response body -> fail open" 0 "this is not json"
pr_json='{"labels": [{"name": "other"}]}'
expect "ready label removed mid-gate -> blocked" 1 "$PC_OK, $DOCS_OK"
exit "$fails"
+22 -2
View File
@@ -1,7 +1,11 @@
name: pre-commit
on:
pull_request:
# pull_request_target instead of pull_request: the workflow definition and
# the hook config are always taken from the BASE branch, so fork /
# first-time-contributor PRs run immediately without a maintainer clicking
# "Approve and run". The PR head is checked out as data only.
pull_request_target:
branches: [main]
workflow_call:
inputs:
@@ -15,12 +19,25 @@ permissions:
jobs:
pre-commit:
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
if: github.event.pull_request.draft != true
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
with:
ref: ${{ inputs.ref || '' }}
# For PR events, lint the PR head — but keep the hook definitions from
# the base branch so an untrusted PR cannot alter what gets executed.
- name: Save trusted hook config
if: github.event_name == 'pull_request_target'
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
- uses: actions/checkout@v4
if: github.event_name == 'pull_request_target'
with:
ref: ${{ github.event.pull_request.head.sha }}
persist-credentials: false
- name: Restore trusted hook config
if: github.event_name == 'pull_request_target'
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
- uses: actions/setup-python@v5
with:
python-version: "3.12"
@@ -30,3 +47,6 @@ jobs:
- uses: pre-commit/action@v3.0.1
with:
extra_args: --all-files --hook-stage manual
# After pre-commit so a self-test failure cannot mask lint failures.
- name: Full-suite gate self-test
run: bash .github/scripts/test_gate_full_suite.sh
+19 -11
View File
@@ -52,6 +52,7 @@ jobs:
core.setOutput('pr_sha', pr.head.sha);
core.setOutput('pr_branch', pr.head.ref);
core.setOutput('pr_number', String(prNumber));
core.setOutput('pr_title', pr.title);
- name: Trigger Full Suite
if: steps.perm.outputs.has_write == 'true'
@@ -60,6 +61,7 @@ jobs:
PR_SHA: ${{ steps.label.outputs.pr_sha }}
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
PR_TITLE: ${{ steps.label.outputs.pr_title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -71,6 +73,7 @@ jobs:
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
--arg pr_title "$PR_TITLE" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
@@ -80,11 +83,12 @@ jobs:
pull_request_id: $pr_id,
pull_request_base_branch: "main",
env: {
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring)
}
}')"
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
}')"
parse-command:
if: >-
@@ -125,7 +129,7 @@ jobs:
set -euo pipefail
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
exit 1
@@ -136,6 +140,7 @@ jobs:
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
[ssim]=ssim [training]=training
[lora-inference]=inference_lora [lora-training]=training_lora
[lora-extraction]=lora_extraction
[distillation]=distillation_dmd [self-forcing]=self_forcing
[vsa]=training_vsa [vmoba]=inference_vmoba
[performance]=performance [api]=api_server
@@ -240,6 +245,7 @@ jobs:
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
PR_TITLE: ${{ github.event.issue.title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -256,6 +262,7 @@ jobs:
--arg full_suite "$FULL_SUITE" \
--arg test_type "$TEST_TYPE" \
--arg pr_number "$PR_NUMBER" \
--arg pr_title "$PR_TITLE" \
'{
commit: $commit,
branch: $branch,
@@ -265,8 +272,9 @@ jobs:
pull_request_base_branch: "main",
env: {
TEST_SCOPE: $test_scope,
FULL_SUITE: $full_suite,
TEST_TYPE: $test_type,
PR_NUMBER: $pr_number
}
}')"
FULL_SUITE: $full_suite,
TEST_TYPE: $test_type,
PR_NUMBER: $pr_number,
PR_TITLE: $pr_title
}
}')"
+21 -1
View File
@@ -7,6 +7,7 @@ on:
permissions:
contents: read
pull-requests: read
actions: read
concurrency:
group: full-suite-${{ github.event.pull_request.number }}
@@ -18,6 +19,8 @@ jobs:
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|| github.event.action == 'synchronize'
runs-on: ubuntu-latest
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
timeout-minutes: 35
steps:
- name: Check ready label
id: check
@@ -49,6 +52,20 @@ jobs:
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
done
# Checks out the BASE branch (default for pull_request_target), so PR
# authors cannot tamper with the gate script.
- name: Checkout gate script
if: steps.check.outputs.has_ready == 'true'
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
- name: Wait for pre-commit and docs build
if: steps.check.outputs.has_ready == 'true'
env:
GH_TOKEN: ${{ github.token }}
PR_SHA: ${{ github.event.pull_request.head.sha }}
PR_NUMBER: ${{ github.event.pull_request.number }}
run: bash .github/scripts/gate_full_suite.sh
- name: Trigger Buildkite Full Suite
if: steps.check.outputs.has_ready == 'true'
env:
@@ -56,6 +73,7 @@ jobs:
PR_SHA: ${{ github.event.pull_request.head.sha }}
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
PR_NUMBER: ${{ github.event.pull_request.number }}
PR_TITLE: ${{ github.event.pull_request.title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -67,6 +85,7 @@ jobs:
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER}" \
--arg pr_title "$PR_TITLE" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
@@ -78,6 +97,7 @@ jobs:
env: {
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring)
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
}')"
+30 -5
View File
@@ -13,12 +13,33 @@ on:
required: false
default: false
type: boolean
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
# push trigger is a sufficient change detector on its own -- no separate
# detect-changes/paths-filter job is needed now that there is a single
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
# rocm Dockerfile stay manual-dispatch only.
push:
branches: [main]
paths:
- 'docker/Dockerfile'
permissions:
contents: read
packages: write
# One static group, no cancellation: every run of this workflow writes the same
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
# concurrent push/dispatch runs would race on those tags, and cancelling a run
# mid-publish can strand the cu126/cu130 tag families at different commits. An
# in-flight superseded build wastes its runner time, but its tags are then
# overwritten by the newer queued run. GitHub keeps a single pending run per
# group: the newest queued run replaces any older queued one.
concurrency:
group: infra-build-image
cancel-in-progress: false
jobs:
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
@@ -28,7 +49,11 @@ jobs:
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
# 2.8.3 comes from the architecture-specific prebuilt releases.
build-cuda-images:
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
# on a push that changed docker/Dockerfile (inputs are null on push). The
# repository guard keeps fork syncs from auto-building; manual dispatch
# still works in forks.
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
strategy:
fail-fast: false
matrix:
@@ -75,10 +100,10 @@ jobs:
secrets: inherit
publish-cuda-manifests:
# !cancelled(): a failed sibling build leg must not skip the manifests for a
# CUDA lane whose own digests all exist; the digest-count check below fails
# the incomplete lane loudly instead.
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
# !cancelled(): publish lanes whose digests exist even if a sibling build
# leg failed (the digest-count check fails incomplete lanes); it also
# bypasses skipped-needs propagation, hence the explicit skipped check.
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
needs: build-cuda-images
runs-on: ubuntu-latest
permissions:
+3 -2
View File
@@ -72,8 +72,7 @@ docs/distillation/examples/
# Python pickle files
*.pkl
# Reference videos
!fastvideo/tests/ssim/reference_videos/**/*.mp4
# Reference videos (negations must come after the catch-all on line below)
# Static images
!docs/assets/images/**/*.png
@@ -127,6 +126,8 @@ apps/dreamverse/web/.env.production.local
.sisyphus/
openspec/
fastvideo/tests/ssim/reference_videos/**
!fastvideo/tests/ssim/reference_videos/**/*.mp4
!fastvideo/tests/ssim/reference_videos/**/*.png
# Editor logs and local Python version pins (accidentally committed)
*.nvimlog
+3
View File
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
# the dreamverse extra and is validated with it, so enable it here.
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
STREAM_MODE=av_fmp4 \
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
FASTVIDEO_FA4=1 \
HF_HOME=/root/.cache/huggingface
RUN mkdir -p /var/lib/dreamverse
+5 -5
View File
@@ -67,7 +67,7 @@ MODEL_REGISTRY = {
},
}
DEFAULT_MODEL_ID = "fast-ltx2"
DEFAULT_MODEL_ID = "fast-ltx23"
ACTIVE_MODEL_ID = (os.getenv("DREAMVERSE_MODEL_ID", "").strip() or DEFAULT_MODEL_ID)
if ACTIVE_MODEL_ID not in MODEL_REGISTRY:
@@ -76,14 +76,11 @@ 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 = 5
NUM_INFERENCE_STEPS = 6
JPEG_QUALITY = 100
BATCH_SIZE = 3
@@ -168,6 +165,9 @@ 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)
+5 -1
View File
@@ -1023,7 +1023,11 @@ 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:
visible_gpu_ids = [int(x.strip()) for x in cuda_visible.split(",") if x.strip()]
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
return _limit_gpu_ids(visible_gpu_ids)
# Auto-detect available GPUs
+3 -1
View File
@@ -206,7 +206,9 @@ def cli() -> None:
args = parser.parse_args()
_install_heartbeat_log_filter()
uvicorn.run(app, host=args.host, port=args.port)
# 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)
if __name__ == "__main__":
+47 -21
View File
@@ -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
from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP, SESSION_TIMEOUT_SECONDS
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,6 +334,7 @@ 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":
@@ -344,7 +345,8 @@ 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(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
generation_paused = bool(not manual_continuation_mode and 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"))
@@ -373,7 +375,10 @@ async def websocket_endpoint(websocket: WebSocket):
prompt_sources_blocked = False
pending_seed_reset = False
pending_seed_reset_reason = ""
pending_simple_submission: PromptSubmission | None = None
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)
single_clip_waiting_for_request = False
rollout_waiting_for_rewrite = False
initial_rollout_waiting_for_rewrite = generation_paused
@@ -393,14 +398,26 @@ 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": 0,
"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,
),
})
if seed_reason == "init":
await ws_send_json({
@@ -493,6 +510,7 @@ 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
@@ -537,6 +555,7 @@ 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)
@@ -544,10 +563,14 @@ async def websocket_endpoint(websocket: WebSocket):
segment_idx = 0
pending_seed_reset = False
pending_seed_reset_reason = ""
pending_simple_submission = None
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)
single_clip_waiting_for_request = False
rollout_waiting_for_rewrite = False
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
generation_paused = bool(not manual_continuation_mode and 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
@@ -968,10 +991,11 @@ 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)
@@ -981,8 +1005,8 @@ async def websocket_endpoint(websocket: WebSocket):
await asyncio.sleep(0.05)
continue
if (not single_clip_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 manual_continuation_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
@@ -1217,7 +1241,9 @@ def cli() -> None:
print(f"Starting mock server with {LATENCY_MS}ms latency on port {args.port}")
_install_heartbeat_log_filter()
uvicorn.run(app, host="0.0.0.0", port=args.port)
# 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)
if __name__ == "__main__":
+60 -18
View File
@@ -313,6 +313,33 @@ 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:
@@ -330,6 +357,7 @@ 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:
@@ -337,21 +365,37 @@ def _parse_json_response(content: str) -> dict[str, Any]:
try:
parsed = json.loads(block)
if isinstance(parsed, dict):
return parsed
last_fenced = parsed
except json.JSONDecodeError:
continue
if last_fenced is not None:
return last_fenced
# Fall back to scanning for the first decodable JSON object in free-form text.
# 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.
decoder = json.JSONDecoder()
for idx, char in enumerate(text):
if char != "{":
continue
last_parsed: dict[str, Any] | None = None
pos = 0
while (idx := text.find("{", pos)) != -1:
try:
parsed, _ = decoder.raw_decode(text[idx:])
parsed, consumed = 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):
return parsed
last_parsed = parsed
if last_parsed is not None:
return last_parsed
raise ValueError("No JSON object found in assistant response.")
@@ -381,7 +425,7 @@ def _format_locked_segments(locked_segments: list[str]) -> str:
class PromptEnhancer:
def __init__(self):
def __init__(self) -> None:
self.provider = PROMPT_PROVIDER
self.provider_label = _resolve_provider_label(PROMPT_PROVIDER)
self.api_key = PROMPT_API_KEY
@@ -1417,15 +1461,12 @@ 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. "
'Respond with valid JSON only as {"next_prompt": "..."}.' # noqa: E501
),
"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."),
}
t0 = time.perf_counter()
@@ -1457,6 +1498,7 @@ 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:
@@ -1539,7 +1581,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 as {"next_prompt": "..."}.' # noqa: E501
'Respond with valid JSON only: {"next_prompt": "<your segment description here>"}.' # noqa: E501
),
}
@@ -20,6 +20,7 @@ from __future__ import annotations
# mypy: ignore-errors
import asyncio
import os
import time
import uuid
from typing import TYPE_CHECKING
@@ -52,6 +53,14 @@ 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."""
@@ -198,6 +207,7 @@ 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()
@@ -282,6 +292,9 @@ 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
@@ -291,13 +304,20 @@ class SessionController:
generation_cap_blocked = False
auto_extension_blocked_segment_idx: int | None = None
prompt_sources_drained_logged = False
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
generation_paused = bool(not manual_continuation_mode and 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
pending_simple_prompt_submission: PromptSubmission | None = None
# 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)
single_clip_waiting_for_request = False
rollout_waiting_for_rewrite = False
initial_rollout_waiting_for_rewrite = generation_paused
@@ -306,6 +326,7 @@ 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
@@ -424,6 +445,7 @@ 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
@@ -458,6 +480,7 @@ 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:
@@ -510,6 +533,7 @@ 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)
@@ -530,10 +554,15 @@ class SessionController:
generation_cap_blocked = False
auto_extension_blocked_segment_idx = None
prompt_sources_drained_logged = False
pending_simple_prompt_submission = None
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)
single_clip_waiting_for_request = False
rollout_waiting_for_rewrite = False
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
generation_paused = bool(not manual_continuation_mode and 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
@@ -942,6 +971,7 @@ class SessionController:
_resolve_generation_segment_cap(
single_clip_mode=single_clip_mode,
cap=GENERATION_SEGMENT_CAP,
manual_continuation_mode=manual_continuation_mode,
),
})
continue
@@ -1005,11 +1035,9 @@ class SessionController:
continue
async def prompt_worker_loop():
while not stop_event.is_set():
try:
submission = await asyncio.wait_for(raw_prompt_queue.get(), timeout=0.1)
except asyncio.TimeoutError:
continue
nonlocal prompt_enhancement_inflight
async def process_submission(submission: PromptSubmission) -> None:
_main_print("INFO", f"Received user prompt for enhancement: {submission.raw_prompt}")
prompt_id = submission.prompt_id
raw_prompt = submission.raw_prompt
@@ -1032,8 +1060,9 @@ class SessionController:
await ws_send_json({
"type": "error",
"message": blocked_raw_prompt_error,
"prompt_id": prompt_id,
})
continue
return
await log_event(
"enhance_request",
{
@@ -1093,8 +1122,9 @@ class SessionController:
await ws_send_json({
"type": "error",
"message": blocked_final_prompt_error,
"prompt_id": prompt_id,
})
continue
return
if result.fallback_used or not final_prompt:
source = "user_enhancement_failed"
_main_print(
@@ -1113,7 +1143,7 @@ class SessionController:
})
# Enhancement is strict JSON-only; do not enqueue raw
# prompt when enhancement fails.
continue
return
else:
source = "user_enhanced"
await ws_send_json({
@@ -1130,6 +1160,7 @@ class SessionController:
source=source,
fallback_used=result.fallback_used,
loop_iteration=loop_iteration,
raw_prompt=raw_prompt,
))
else:
await ready_prompt_queue.put(
@@ -1139,6 +1170,7 @@ class SessionController:
source="user_raw",
fallback_used=False,
loop_iteration=loop_iteration,
raw_prompt=raw_prompt,
))
await ws_send_json({
"type": "prompt_ready",
@@ -1148,6 +1180,21 @@ 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(),
@@ -1300,6 +1347,7 @@ 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({
@@ -1359,6 +1407,7 @@ 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":
@@ -1388,8 +1437,9 @@ class SessionController:
await raw_prompt_queue.put(pending_simple_prompt_submission)
pending_simple_prompt_submission = None
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):
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):
loop_generation_enabled = False
rollout_waiting_for_rewrite = True
_main_print(
@@ -1542,7 +1592,10 @@ class SessionController:
if single_clip_mode:
await asyncio.sleep(PROMPT_AUTO_SLEEP_MS / 1000.0)
continue
if not prompt_sources_drained_logged:
# 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:
snapshot = queue_snapshot()
_main_print(
"WARN",
@@ -1581,6 +1634,24 @@ 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,3 +19,4 @@ class ReadyPrompt:
fallback_used: bool = False
seed_prompt_index: int | None = None
loop_iteration: int | None = None
raw_prompt: str | None = None
@@ -150,12 +150,22 @@ 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,12 +75,13 @@ 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) -> None:
def run(app, host: str, port: int, **kwargs) -> None:
calls.append(
{
"app": app,
"host": host,
"port": port,
**kwargs,
}
)
@@ -104,6 +105,7 @@ 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,
}
]
@@ -121,6 +123,7 @@ 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,
}
]
@@ -147,6 +150,7 @@ 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,
}
]
@@ -166,6 +170,7 @@ 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,6 +290,152 @@ 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,6 +6,7 @@ import os
import re
import time
import pytest
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -205,6 +206,59 @@ 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")
+3 -2
View File
@@ -15,5 +15,6 @@ 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) -> int:
return 0 if single_clip_mode else cap
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
@@ -489,7 +489,7 @@ class VideoGenerationWorker:
num_inference_steps=NUM_INFERENCE_STEPS,
guidance_scale=1.0,
seed=10,
ltx2_image_crf=0.0,
ltx2_image_crf=(33.0 if image_path and segment_idx == 1 else 0.0),
image_path=image_path if segment_idx == 1 else None,
return_continuation_state=False,
)
+3
View File
@@ -29,6 +29,9 @@ test = [
dreamverse-server = "dreamverse.server_entry:cli"
dreamverse-mock-server = "dreamverse.mock_server:cli"
[tool.setuptools.packages.find]
include = ["dreamverse*"]
[tool.uv]
package = false
+67
View File
@@ -0,0 +1,67 @@
#!/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"
+30
View File
@@ -0,0 +1,30 @@
#!/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
+105
View File
@@ -0,0 +1,105 @@
#!/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:33 SESSION_TIMEOUT_SECONDS = 300
# internal: config.py SESSION_TIMEOUT_SECONDS (env DREAMVERSE_SESSION_TIMEOUT_SECONDS, default 300)
session_timeout_seconds: 300
# internal: config.py:282-284 GENERATION_SEGMENT_CAP default 6
generation_segment_cap: 6
+125 -13
View File
@@ -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 from "@/components/Workspace";
import Workspace, { SceneHistoryList } 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 } from "@/stores/stream";
import { createStreamStore, USER_PROMPT_SOURCES } 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";
const HERO_TEXT = "Direct scenes in seconds with";
function HeroTagline() {
const ref = useRef<HTMLHeadingElement>(null);
@@ -166,6 +166,8 @@ 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>
);
}
@@ -237,6 +239,7 @@ export default function Page() {
enhancementEnabled,
promptExtensionError,
autoExtensionEnabled,
manualContinuationMode,
autoExtensionTimeoutHint,
loopGenerationEnabled,
generationPaused,
@@ -249,6 +252,8 @@ export default function Page() {
livePromptRewriteMode,
sessionExpired,
projectResetPending,
waitingForSegmentPrompt,
generatingNextScene,
} = sessionState;
const {
@@ -323,6 +328,11 @@ 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);
@@ -422,6 +432,11 @@ 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]);
@@ -432,11 +447,39 @@ export default function Page() {
const canDownloadVideo = useMemo(() => {
const currentActiveClip = activeClip as Record<string, any> | null;
if (currentActiveClip?.blob instanceof Blob) return true;
return (completedClips as Record<string, any>[]).some((clip) => clip?.blob instanceof Blob);
}, [activeClip, completedClips]);
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]);
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],
);
@@ -1421,7 +1464,8 @@ export default function Page() {
if (!prompt) return;
lastSubmitTimeRef.current = now;
const rCAR = !uiStore.get().devtoolsMode && !uiStore.get().demoMode;
if (rCAR || (sessionStore.get().livePromptRewriteMode && !uiStore.get().demoMode)) {
const inManualMode = sessionStore.get().manualContinuationMode;
if (!inManualMode && (rCAR || (sessionStore.get().livePromptRewriteMode && !uiStore.get().demoMode))) {
if (rewriteStore.get().rewritingSeedPrompts) return;
const rewriteSourcePromptWindowPrompts = getActivePromptWindowPrompts();
const nextPendingClip = {
@@ -1462,6 +1506,11 @@ 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({
@@ -1475,7 +1524,14 @@ export default function Page() {
activeClipId: shouldUseArchivedPlaybackFallback() ? streamStore.get().activeClipId : "",
activePlaybackStartTime: shouldUseArchivedPlaybackFallback() ? streamStore.get().activePlaybackStartTime : 0,
});
sessionStore.patch({ livePromptDraft: "" });
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 } : {}),
});
}
function setLivePromptRewriteMode(enabled: boolean) {
@@ -1550,6 +1606,14 @@ 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),
@@ -1725,6 +1789,9 @@ 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();
@@ -1737,6 +1804,7 @@ export default function Page() {
function resetToProjectLobbyState() {
setVideoMuted(true);
pendingInitialPromptRef.current = "";
pendingInitialPromptIdRef.current = "";
sessionStore.patch({
sessionStarted: false,
livePromptDraft: "",
@@ -1750,6 +1818,9 @@ export default function Page() {
sessionNotice: "",
sessionExpired: false,
projectResetPending: false,
manualContinuationMode: true,
waitingForSegmentPrompt: false,
generatingNextScene: false,
});
rewriteStore.resetSessionState();
streamStore.resetSessionState();
@@ -1760,7 +1831,14 @@ export default function Page() {
}
function buildProjectInitPayload(type: "session_init_v2" | "project_init_v1") {
const segmentPrompts = getSessionInitPrompts();
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);
}
setSeedPrompts(segmentPrompts);
return {
type,
@@ -1768,11 +1846,17 @@ export default function Page() {
preset_label: getInitialPresetLabel(),
curated_prompts: segmentPrompts,
initial_rollout_prompt: normalizeInitialPrompt(pendingInitialPromptRef.current),
initial_image: null,
// 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,
single_clip_mode: false,
enhancement_enabled: sessionStore.get().enhancementEnabled,
auto_extension_enabled: sessionStore.get().autoExtensionEnabled,
loop_generation_enabled: sessionStore.get().loopGenerationEnabled,
auto_extension_enabled: manualMode ? false : sessionStore.get().autoExtensionEnabled,
loop_generation_enabled: manualMode ? false : sessionStore.get().loopGenerationEnabled,
manual_continuation_mode: manualMode,
};
}
@@ -1951,7 +2035,11 @@ 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,
@@ -1966,6 +2054,9 @@ export default function Page() {
autoExtensionTimeoutHint: "",
generationPaused: false,
projectResetPending: false,
manualContinuationMode: nextManualContinuationMode,
waitingForSegmentPrompt: false,
generatingNextScene: false,
});
resetPlaybackState();
streamStore.patch({
@@ -1979,12 +2070,15 @@ export default function Page() {
selectedHistoryId: "",
});
rewriteStore.resetSessionState();
pendingInitialPromptIdRef.current = initialPrompt ? makePromptId() : "";
if (initialPrompt) {
addPromptEvent({
promptId: makePromptId(),
promptId: pendingInitialPromptIdRef.current,
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());
@@ -2500,6 +2594,7 @@ 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}
@@ -2512,6 +2607,7 @@ export default function Page() {
onEnhancementToggle={handleEnhancementToggle}
onCuratedPromptLimitChange={handleCuratedPromptLimitChange}
onAutoExtensionToggle={handleAutoExtensionToggle}
onManualContinuationToggle={handleManualContinuationToggle}
onLoopToggle={handleLoopGenerationToggle}
onJoin={joinSession}
onLeave={leaveSession}
@@ -2641,6 +2737,7 @@ 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 && (
<>
@@ -2724,6 +2821,8 @@ export default function Page() {
showLivePlayback={showLivePlayback}
defaultMuted={videoMuted}
canDownload={canDownloadVideo}
waitingForSegmentPrompt={waitingForSegmentPrompt as boolean}
generatingNextScene={generatingNextScene as boolean}
onPlaying={markFirstFrameRendered}
onDownload={handleDownloadVideo}
/>
@@ -2734,6 +2833,7 @@ 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}
@@ -2780,11 +2880,17 @@ 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}
@@ -2798,6 +2904,12 @@ export default function Page() {
</motion.div>
</div>
</div>
{manualContinuationMode && (
<div className="px-4 sm:px-6">
<SceneHistoryList sceneHistory={steeringScenes as any[]} />
</div>
)}
</div>
</main>
);
}
+96 -6
View File
@@ -2,13 +2,16 @@
import React, { useRef, useState, useCallback, useEffect } from "react";
import Image from "next/image";
import { Film, ArrowUp, X, Loader2, ArrowLeft } from "lucide-react";
import { Film, ArrowUp, X, Loader2, ArrowLeft, ImagePlus } 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;
@@ -22,6 +25,12 @@ 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;
@@ -46,6 +55,12 @@ export default function ChatBar({
sessionNotice = "",
projectResetPending = false,
viewingReadOnly = false,
waitingForSegmentPrompt = false,
manualContinuationEnabled = false,
onModeChange = () => {},
initialImageDataUrl = "",
onImageUpload = () => {},
onImageClear = () => {},
onPresetGenerate = () => {},
onContinuationInput = () => {},
onContinuationKeydown = () => {},
@@ -59,15 +74,51 @@ 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"
: !sessionStarted
? "What video are you imagining?"
: "What do you want to edit?";
: waitingForSegmentPrompt
? "Describe the next scene\u2026"
: !sessionStarted
? "What video are you imagining?"
: "What do you want to edit?";
const actionLabel = !sessionStarted ? "Generate" : "Rewrite rollout";
const inputRef = useRef<HTMLTextAreaElement>(null);
@@ -267,9 +318,9 @@ export default function ChatBar({
<Button onClick={onStartNewProject} size="sm" className="rounded-full px-5">
New Project
</Button>
<a href="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
<a href="https://haoailab.com/blogs/dreamverse/" target="_blank" rel="noopener noreferrer">
<Button variant="outline" size="sm" className="rounded-full px-5">
Join Waitlist
Blog
</Button>
</a>
</div>
@@ -346,6 +397,31 @@ 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",
@@ -359,6 +435,7 @@ export default function ChatBar({
value={continuationDraft}
onChange={onContinuationInput}
onKeyDown={handleKeyDown}
onPaste={handleImagePaste}
placeholder={sttBusy ? "Listening\u2026" : messagePlaceholder}
maxLength={PROMPT_MAX_LENGTH}
disabled={isBusy || sttBusy}
@@ -369,6 +446,19 @@ 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,7 +8,6 @@ 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 {
@@ -29,13 +28,13 @@ export default function Header({ timeLeft = null, formatTime = (seconds) => `${s
<SidePanelOpenFilled size={20} />
</Button>
)}
<a href={FASTVIDEO_REPO_URL} target="_blank" rel="noopener noreferrer" title="FastVideo on GitHub">
<a href="/" title="FastVideo home">
<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="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
<a href={FASTVIDEO_BLOG_URL} target="_blank" rel="noopener noreferrer">
<Button variant="outline" size="sm" className="gap-1.5 rounded-full px-3 text-xs">
Join Waitlist
Blog
<ExternalLink className="size-3 opacity-60" />
</Button>
</a>
@@ -53,9 +52,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="https://docs.google.com/forms/d/e/1FAIpQLSe5zpO1iD8Ds-Ih-fOLm64qd7YZVvuvAyHuJaAfw1hkRHTe_A/viewform?usp=publish-editor" target="_blank" rel="noopener noreferrer">
<a href={FASTVIDEO_BLOG_URL} target="_blank" rel="noopener noreferrer">
<Button variant="outline" size="sm" className="gap-1.5 rounded-full px-3 text-xs">
Join Waitlist
Blog
<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 { Download, Loader2, Share } from "lucide-react";
import { Check, ChevronDown, Download, Loader2, Share } from "lucide-react";
import { Button } from "@/components/ui/button";
interface VideoPlayerProps {
videoRef?: React.RefCallback<HTMLVideoElement>;
@@ -21,6 +21,8 @@ interface VideoPlayerProps {
showLivePlayback?: boolean;
defaultMuted?: boolean;
rewritePending?: boolean;
waitingForSegmentPrompt?: boolean;
generatingNextScene?: boolean;
onPlaying?: () => void;
onDownload?: () => void;
}
@@ -57,6 +59,8 @@ export default function VideoPlayer({
showLivePlayback = true,
defaultMuted = true,
rewritePending = false,
waitingForSegmentPrompt = false,
generatingNextScene = false,
onPlaying = () => {},
onDownload,
}: VideoPlayerProps) {
@@ -88,6 +92,137 @@ 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">
@@ -104,7 +239,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 && loadingAnimation ? (
) : !avPlaybackStarted && !mediaAppendError && !inQueue && !waitingForSegmentPrompt && 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" />
@@ -114,6 +249,35 @@ 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&hellip;</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,11 +3,90 @@ 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, Lightbulb, Pencil } from "lucide-react";
import { Check, Clapperboard, 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;
@@ -19,6 +98,7 @@ interface WorkspaceProps {
selectedClipId?: string;
selectedEntryKey?: string;
originalClipId?: string;
manualMode?: boolean;
}
function normalizeText(value: any): string {
@@ -218,7 +298,7 @@ function ChromaGradient({ sessionStarted = false }: { sessionStarted?: boolean }
);
}
export default function Workspace({ promptEvents = [], currentThumbnail = null, originalLabel = "", sessionStarted = false, onSelectOriginal, onSelectEvent, onSelectCurrent, selectedClipId, selectedEntryKey: selectedEntryKeyProp, originalClipId = "" }: WorkspaceProps) {
export default function Workspace({ promptEvents = [], currentThumbnail = null, originalLabel = "", sessionStarted = false, onSelectOriginal, onSelectEvent, onSelectCurrent, selectedClipId, selectedEntryKey: selectedEntryKeyProp, originalClipId = "", manualMode = false }: WorkspaceProps) {
const bottomSentinelRef = useRef<HTMLDivElement>(null);
const topSentinelRef = useRef<HTMLDivElement>(null);
const [showTopFade, setShowTopFade] = useState(false);
@@ -259,6 +339,14 @@ 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,6 +34,7 @@ interface DevtoolsComposerProps {
demoMode?: boolean;
enhancementEnabled?: boolean;
autoExtensionEnabled?: boolean;
manualContinuationEnabled?: boolean;
loopGenerationEnabled?: boolean;
curatedPromptLimit?: number;
maxCuratedPromptCount?: number;
@@ -49,6 +50,7 @@ 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;
@@ -71,6 +73,7 @@ export default function DevtoolsComposer({
demoMode = false,
enhancementEnabled = true,
autoExtensionEnabled = false,
manualContinuationEnabled = false,
loopGenerationEnabled = false,
curatedPromptLimit = 0,
maxCuratedPromptCount = 0,
@@ -86,6 +89,7 @@ export default function DevtoolsComposer({
onEnhancementToggle = () => {},
onCuratedPromptLimitChange = () => {},
onAutoExtensionToggle = () => {},
onManualContinuationToggle = () => {},
onLoopToggle = () => {},
onLivePromptModeToggle = () => {},
onSpeechTranscript,
@@ -328,6 +332,28 @@ 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,6 +21,7 @@ interface DevtoolsShellProps {
selectedPresetId?: string;
enhancementEnabled?: boolean;
autoExtensionEnabled?: boolean;
manualContinuationEnabled?: boolean;
loopGenerationEnabled?: boolean;
canJoinSession?: boolean;
canSubmitContinuation?: boolean;
@@ -34,6 +35,7 @@ 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;
@@ -128,6 +130,7 @@ export default function DevtoolsShell({
selectedPresetId = '',
enhancementEnabled = true,
autoExtensionEnabled = false,
manualContinuationEnabled = false,
loopGenerationEnabled = false,
canJoinSession = false,
canSubmitContinuation = false,
@@ -141,6 +144,7 @@ export default function DevtoolsShell({
onEnhancementToggle = () => {},
onCuratedPromptLimitChange = () => {},
onAutoExtensionToggle = () => {},
onManualContinuationToggle = () => {},
onLoopToggle = () => {},
onJoin = () => {},
onLeave = () => {},
@@ -284,6 +288,7 @@ export default function DevtoolsShell({
demoMode={demoMode}
enhancementEnabled={enhancementEnabled}
autoExtensionEnabled={autoExtensionEnabled}
manualContinuationEnabled={manualContinuationEnabled}
loopGenerationEnabled={loopGenerationEnabled}
curatedPromptLimit={curatedPromptLimit}
maxCuratedPromptCount={maxCuratedPromptCount}
@@ -299,6 +304,7 @@ export default function DevtoolsShell({
onEnhancementToggle={onEnhancementToggle}
onCuratedPromptLimitChange={onCuratedPromptLimitChange}
onAutoExtensionToggle={onAutoExtensionToggle}
onManualContinuationToggle={onManualContinuationToggle}
onLoopToggle={onLoopToggle}
onLivePromptModeToggle={onLivePromptModeToggle}
onSpeechTranscript={onSpeechTranscript}
@@ -57,4 +57,32 @@ 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');
});
});
+15 -1
View File
@@ -16,5 +16,19 @@ export function prependPromptEvent(
events: Record<string, any>[],
event: Record<string, any>,
): Record<string, any>[] {
return [event, ...events].slice(0, MAX_PROMPT_EVENTS);
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;
});
}
+165 -1
View File
@@ -1,6 +1,11 @@
import { describe, expect, it } from 'vitest';
import { resolveSessionErrorMessage } from './reducer';
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';
describe('resolveSessionErrorMessage', () => {
it('returns a dedicated message for IP session limit errors', () => {
@@ -19,3 +24,162 @@ 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);
});
});
+50 -9
View File
@@ -99,7 +99,19 @@ 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;
@@ -243,19 +255,32 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
}
case "prompt/sources_blocked":
sessionStore.patch({
autoExtensionTimeoutHint: uiStore.get().simpleMode ? "" : "blocked on user input, increase prompt count for smoother experience",
});
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",
});
}
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":
sessionStore.patch({ autoExtensionTimeoutHint: "" });
if (event.type === "session/auto_extension_updated") {
console.log("[AutoExtensionUpdated]", {
enabled: sessionStore.get().autoExtensionEnabled,
});
}
// Deliberately does NOT touch generatingNextScene: toggling auto extension
// starts no generation.
sessionStore.patch({ autoExtensionTimeoutHint: "", waitingForSegmentPrompt: false });
console.log("[AutoExtensionUpdated]", {
enabled: sessionStore.get().autoExtensionEnabled,
});
return;
case "segment/step_complete":
@@ -277,6 +302,7 @@ export async function applyNormalizedSocketEvent(event: any, context: any): Prom
projectResetPending: false,
sessionExpired: true,
sessionNotice: "",
generatingNextScene: false,
});
console.log("Session timed out");
clearCountdownInterval();
@@ -354,6 +380,8 @@ 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,
@@ -471,17 +499,30 @@ 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,6 +24,9 @@ export interface SessionState {
livePromptRewriteMode: boolean;
sessionExpired: boolean;
projectResetPending: boolean;
manualContinuationMode: boolean;
waitingForSegmentPrompt: boolean;
generatingNextScene: boolean;
}
const DEFAULT_SESSION_STATE: SessionState = {
@@ -49,6 +52,12 @@ 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> & {
+18 -1
View File
@@ -1,5 +1,18 @@
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;
@@ -159,10 +172,14 @@ export function createStreamStore(initialState: Partial<StreamState> = {}): Stre
loopIteration: typeof loopIteration === "number" ? loopIteration : null,
};
const nextHistory = [entry, ...state.promptHistory];
return {
...state,
promptHistoryCounter: nextCounter,
promptHistory: [entry, ...state.promptHistory].slice(0, 120),
promptHistory: nextHistory.length > PROMPT_HISTORY_CAP
? nextHistory.filter((item, index) =>
index < PROMPT_HISTORY_CAP || !USER_PROMPT_SOURCES.has(String(item.source || "")))
: nextHistory,
selectedHistoryId: state.selectedHistoryId || (entry.id as string),
};
});
+16 -6
View File
@@ -12,13 +12,17 @@ Defaults:
- `HF_REPO_ID=FastVideo/performance-tracking`
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
- `PERF_MAX_REGRESSION=0.05`
Records can include source metadata:
Records can include source metadata and rolling-baseline policy context:
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
- `baseline_eligible`: only successful scheduled-main records should be true
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
floors used for recomputed status context
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
crossings; `regressed` remains the gated CI-failure signal.
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
configured dataset repo requires authenticated access:
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
- PR number, branch, and Buildkite URL when present
The latest status table uses the stored JSON `success` value. Recomputed
baseline context is shown separately and does not override stored status.
baseline context applies each metric's percent and absolute regression floors
and does not override stored status.
## API
@@ -99,6 +104,11 @@ baseline context is shown separately and does not override stored status.
- `GET /api/performance/trends?days=90&run_source=scheduled_main`
- `GET /api/performance/records?days=90&run_source=local`
The current v1 grouping key is `(model_id, gpu_type)`. Baselines are computed
from the latest five previous successful records in each group for dashboard
context. CI gating uses only records marked `baseline_eligible=true`.
V2 records use the same comparison cohort as CI: `workload_id`, `variant_id`,
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
`software_profile_id`. `model_id` and `gpu_type` remain display/filter
metadata, so renaming either does not split history. Legacy records still group
by `(model_id, gpu_type)`. Dashboard baselines use the latest five previous
successful, baseline-eligible records in each group. Summary and trend filters
match the latest display metadata after grouping, while the raw records endpoint
continues to filter individual records.
@@ -1,6 +1,7 @@
import { useEffect, useMemo, useState } from "react";
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
import { fetchSummary, fetchTrends, refreshData } from "./api";
import type { CohortValue, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
@@ -109,6 +110,61 @@ function metricLabel(metricKey: string) {
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
}
type CohortFields = {
model_id: string;
gpu_type: string;
workload_id: CohortValue;
variant_id: CohortValue;
benchmark_version: CohortValue;
recipe_fingerprint: CohortValue;
hardware_profile_id: CohortValue;
software_profile_id: CohortValue;
};
function cohortValue(value: CohortValue) {
if (value === null || value === undefined || value === "") {
return "legacy";
}
return String(value);
}
function shortCohortValue(value: CohortValue) {
const text = cohortValue(value);
if (text === "legacy" || text.length <= 14) {
return text;
}
return text.slice(0, 12);
}
function cohortKey(cohort: CohortFields) {
return [
cohort.model_id,
cohort.gpu_type,
cohortValue(cohort.workload_id),
cohortValue(cohort.variant_id),
cohortValue(cohort.benchmark_version),
cohortValue(cohort.recipe_fingerprint),
cohortValue(cohort.hardware_profile_id),
cohortValue(cohort.software_profile_id)
].join("|");
}
function cohortTitle(cohort: CohortFields) {
const workload = cohortValue(cohort.workload_id);
const variant = cohortValue(cohort.variant_id);
const version = cohortValue(cohort.benchmark_version);
const versionLabel = version === "legacy" ? version : `v${version}`;
return `${workload} / ${variant} / ${versionLabel}`;
}
function cohortDetail(cohort: CohortFields) {
return [
`recipe ${shortCohortValue(cohort.recipe_fingerprint)}`,
shortCohortValue(cohort.hardware_profile_id),
shortCohortValue(cohort.software_profile_id)
].join(" | ");
}
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
const definition = METRIC_DEFINITIONS[metricKey];
if (!definition) {
@@ -171,7 +227,9 @@ function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string
top: `${(activePoint.y / height) * 100}%`
}
: undefined;
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}, ${cohortTitle(
group
)}`;
return (
<div className="chart-shell">
@@ -419,7 +477,7 @@ export default function App() {
<section className="panel">
<div className="panel-header">
<h2>Latest Status</h2>
<span>{latestRows.length} model/GPU groups</span>
<span>{latestRows.length} comparison cohorts</span>
</div>
{latestRows.length === 0 ? (
<div className="empty">No records match the selected filters.</div>
@@ -432,6 +490,7 @@ export default function App() {
<th>Recomputed</th>
<th>Model</th>
<th>GPU</th>
<th>Cohort</th>
<th>Commit</th>
<th>Source</th>
<th>Baseline</th>
@@ -440,11 +499,13 @@ export default function App() {
<th>Throughput</th>
<th>Memory</th>
<th>Worst</th>
<th>Exceeded</th>
<th>Failing</th>
</tr>
</thead>
<tbody>
{latestRows.map((row) => (
<tr key={`${row.model_id}-${row.gpu_type}`}>
<tr key={cohortKey(row)}>
<td>
<span className={`badge ${row.status}`}>{row.status}</span>
</td>
@@ -455,6 +516,12 @@ export default function App() {
</td>
<td>{row.model_id}</td>
<td>{row.gpu_type}</td>
<td>
<div className="cohort-cell">
<strong>{cohortTitle(row)}</strong>
<span>{cohortDetail(row)}</span>
</div>
</td>
<td>{shortSha(row.commit_sha)}</td>
<td>
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
@@ -465,6 +532,12 @@ export default function App() {
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
<td>
{row.threshold_exceeded_metrics.length
? row.threshold_exceeded_metrics.join(", ")
: "none"}
</td>
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
</tr>
))}
</tbody>
@@ -487,11 +560,13 @@ export default function App() {
) : (
trends.map((group) =>
METRIC_KEYS.map((metricKey) => (
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
<article className="trend-card" key={`${cohortKey(group)}-${metricKey}`}>
<div>
<h3>{metricLabel(metricKey)}</h3>
<p>
{group.model_id} | {group.gpu_type}
<span>{cohortTitle(group)}</span>
<span>{cohortDetail(group)}</span>
</p>
</div>
<TrendChart group={group} metricKey={metricKey} />
+22 -4
View File
@@ -2,11 +2,28 @@ export type MetricValue = {
current: number | null;
baseline: number | null;
regression_pct: number | null;
absolute_delta: number | null;
threshold_percent: number;
threshold_absolute: number;
gated: boolean;
threshold_exceeded: boolean;
regressed: boolean;
label: string;
lower_is_better: boolean;
precision: number;
};
export type CohortValue = string | number | null;
export type ComparisonCohort = {
workload_id: CohortValue;
variant_id: CohortValue;
benchmark_version: CohortValue;
recipe_fingerprint: CohortValue;
hardware_profile_id: CohortValue;
software_profile_id: CohortValue;
};
export type SummaryRow = {
model_id: string;
gpu_type: string;
@@ -15,7 +32,8 @@ export type SummaryRow = {
success: boolean;
baseline_n: number;
worst_regression_pct: number | null;
regression_threshold_pct: number;
threshold_exceeded_metrics: string[];
failing_metrics: string[];
computed_regression_status: "pass" | "fail";
status: "pass" | "fail";
run_source: RunSource;
@@ -27,7 +45,7 @@ export type SummaryRow = {
build_id: string;
job_id: string;
metrics: Record<string, MetricValue>;
};
} & ComparisonCohort;
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
@@ -61,13 +79,13 @@ export type TrendPoint = {
build_id: string;
job_id: string;
metrics: Record<string, number | null>;
};
} & ComparisonCohort;
export type TrendGroup = {
model_id: string;
gpu_type: string;
points: TrendPoint[];
};
} & ComparisonCohort;
export type TrendsResponse = {
groups: TrendGroup[];
@@ -149,6 +149,11 @@ h3 {
font-size: 0.82rem;
}
.trend-card p {
display: grid;
gap: 2px;
}
.stat strong {
display: block;
margin-top: 8px;
@@ -186,7 +191,7 @@ h3 {
table {
width: 100%;
min-width: 1120px;
min-width: 1260px;
border-collapse: collapse;
}
@@ -209,6 +214,25 @@ td {
font-size: 0.9rem;
}
.cohort-cell {
display: grid;
gap: 2px;
}
.cohort-cell strong,
.trend-card p span {
color: #1b2836;
font-size: 0.78rem;
font-weight: 700;
}
.cohort-cell span,
.trend-card p span + span {
color: #607080;
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", monospace;
font-size: 0.72rem;
}
.badge {
display: inline-flex;
align-items: center;
@@ -0,0 +1,22 @@
{
"alpha_yaw": 0.08734091699186919,
"alpha_pitch": 0.08169667696275307,
"alpha_turn": 5.724587470723463e-17,
"beta_fwd": 0.02842768078408099,
"beta_strafe": 0.022531015077067108,
"focal_length": 457.0,
"frame_shape": [
352,
640
],
"calibrated_from": [
"1_wasd_only",
"camera",
"camera4hold_alpha1",
"fully_random",
"wasdonly_alpha1",
"wasd4holdrandview_simple_1key1mouse1"
],
"residual_rms": 15.890399609478676,
"n_equations": 4125232
}
+6 -4
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=940cd9680f3315f2f06b43ab5bea2c2cf2d96806
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
@@ -161,7 +161,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
uv pip install flash-attn==${FLASH_ATTN_VERSION} --no-build-isolation; \
fi
# Overlay the cutlass-4.5-safe upstream FA4 cute (FA4_CUTE_REF) over the
# Overlay the CuTe-DSL-4.6-compatible upstream FA4 cute (FA4_CUTE_REF) over the
# wheel/source one so the image runs FA4, not the FA2 fallback. This pulls the FA4
# runtime stack (cutlass-dsl, quack-kernels, apache-tvm-ffi, torch-c-dlpack-ext) --
# the same deps the [dreamverse] extra already installs in CI; the installed torch
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
# Then verify both survive so a broken overlay fails the build instead of shipping
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
# arm64 and the image runs FA3/FA2 as usual.
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
if [ "${TARGETARCH}" = "arm64" ]; then \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
else \
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
+9
View File
@@ -99,10 +99,18 @@ status.
Full Suite is also path-filtered. It validates broader behavior before Mergify
can merge a PR.
A `ready`-labeled PR does not hit Buildkite immediately:
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
which waits for the cheap Tier-1 checks (pre-commit, docs build) on the PR
head. A red cheap check blocks the suite (fail closed; the next push re-arms
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
`/test full` bypasses the gate.
| Buildkite label | `TEST_TYPE` | Main watched paths |
|---|---|---|
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
@@ -144,6 +152,7 @@ Valid direct test names:
| `/test training` | `training` |
| `/test lora-inference` | `inference_lora` |
| `/test lora-training` | `training_lora` |
| `/test lora-extraction` | `lora_extraction` |
| `/test distillation` | `distillation_dmd` |
| `/test self-forcing` | `self_forcing` |
| `/test vsa` | `training_vsa` |
+299 -53
View File
@@ -12,7 +12,8 @@ It serves three audiences:
* **Maintainers** — surfaces regressions in a Markdown summary on every
performance build and a long-form Plotly dashboard.
* **Local developers** — lets you run the same benchmark on your own machine,
then compare against the historical baseline for the same model and GPU.
then compare against the historical baseline for the same comparable
identity.
## Quick start (local)
@@ -72,14 +73,18 @@ fastvideo/tests/performance/
│ writes Markdown summary + (optionally) uploads new records
├── dashboard.py
│ └── builds time-series Plotly HTML from HF history
└── hf_store.py # shared HF I/O + DataFrame helpers
fastvideo/performance/
├── hf_store.py # shared HF I/O + DataFrame helpers
└── metric_policy.py # shared rolling-baseline threshold policy
```
The HF dataset (`FastVideo/performance-tracking` by default) holds one
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
the median of the last 5 successful, baseline-eligible records for that
model+GPU. PR and local records are visible in the dashboard but are not
baseline eligible.
normalized JSON per run. For v2 records, the rolling baseline is the median of
the last 5 successful, baseline-eligible records in the same comparison cohort:
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
`hardware_profile_id`, and `software_profile_id`. PR and local records are
visible in the dashboard but are not baseline eligible.
## Planned Coverage
@@ -92,25 +97,28 @@ and recipe changes instead of treating all records for a model as equivalent.
## Metrics
Each benchmark records six metrics:
Each benchmark records six metrics. The rolling-baseline comparator also has a
per-metric policy with direction, percent threshold, absolute threshold, and a
`gated` flag.
| Metric | Raw key | Normalized key | Direction |
|---|---|---|---|
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|---|---|---|---|---|
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
while it runs so pipeline stage execution times are available in
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
`prompt_encoding_stage` so duplicate stage classes do not collide. For
`PipelineStage` entries, the extractor maps the `stage_class` field:
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
`vae_decode_time_s`, with a fallback for older logs that used the class name as
`PipelineStage` entries, shared component stage bases emit a stable
`component_metric`: text encoding stages map to `text_encoder_time_s`,
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
older logs that do not include `component_metric` or that used the class name as
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
does not report one of the mapped stages, that component metric is stored as
@@ -152,26 +160,115 @@ unrealistic memory growth, and optionally large component-specific slowdowns
even when the rolling baseline is empty. They are hand-set with generous
headroom and almost never need touching.
### Rolling baseline (per `(model_id, gpu_type)`)
### Rolling baseline (per comparison cohort)
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
for each available metric, and fails if the current run regresses by more than
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
higher values are regressions. For throughput, lower values are regressions.
for the same comparison cohort from the HF dataset, computes the median for
each available metric, and evaluates the current run with the metric's rolling
regression policy. For v2 records, that cohort is `workload_id`, `variant_id`,
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
`software_profile_id`. For latency, memory, and component times, higher values
are regressions. For throughput, lower values are regressions.
A metric exceeds its rolling threshold when both of these are true:
```text
percent_delta > threshold_percent
absolute_delta > threshold_absolute
```
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
for metrics that should remain visible in reports and the dashboard without
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
metrics are skipped.
This is the **drift detector** — it catches sub-threshold regressions that
slowly add up. Only scheduled-main successful records are baseline eligible.
Local and pull-request runs can upload dashboard-visible records, but they do
not update future gating baselines.
Comparator summaries and normalized artifacts include an explicit
`comparison_status`:
| Status | Meaning | CI behavior |
|---|---|---|
| `PASS` | Comparable baseline exists and no gated metric regressed. Legacy records with no baseline also keep the historical initialization behavior. | Passes |
| `REGRESSION` | The record exceeds one of its static thresholds or at least one gated metric regressed against a comparable baseline. | Fails |
| `CALIBRATION_NEEDED` | A v2 record has no exact comparable baseline. | Passes, may upload when `PERF_UPLOAD_POLICY=pass`, never seeds a baseline |
| `RECIPE_MISMATCH` | The same workload, variant, and benchmark version has trusted successful records under another recipe fingerprint, including records from other hardware or software profiles. | Fails |
| `INFRA_ERROR` | The comparator cannot safely classify the record, such as a v2 record missing required identity fields. | Fails |
`QUALITY_BLOCKED` is reserved for a future variant-promotion workflow and is
not emitted by normal rolling-baseline comparison.
For `RECIPE_MISMATCH`, trusted records are scheduled-main records or records
already marked `baseline_eligible`. PR and local calibration uploads remain
visible but do not authorize future CI-gating recipe mismatch failures.
When the baseline shifts for a legitimate reason (torch upgrade, kernel
change, etc.) and CI starts failing, use the
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
agent skill to advance the rolling median.
agent skill to advance the rolling median. To approve the first baseline for
a new v2 exact identity, follow that skill with a reviewed scheduled-main
full-suite `CALIBRATION_NEEDED` normalized artifact. Its prepare step uses
`fastvideo/tests/performance/seed_baseline.py`; after a separate human gate,
the skill rechecks the current HF revision and conditionally uploads the whole
seed batch in one commit.
## Schemas
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
Benchmark configs without `config_schema_version` are treated as legacy v1
configs and remain loadable. New or migrated configs should use
`config_schema_version: 2` and include explicit comparable identity fields:
```jsonc
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 3
}
```
`benchmark_id` is still required because raw artifact names, generated-video
directories, normalized record paths, and legacy storage directories depend on
it. The v2 comparator does not use it as part of the comparison cohort. The v2
identity fields make the measured workload explicit:
| Field | Purpose |
|---|---|
| `workload_id` | Stable benchmark family, such as `wan-t2v`. |
| `variant_id` | Intentional recipe family, including model size and parallelism config, such as `1.3b-sp2`. |
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
If a config declares `config_schema_version: 2`, loading fails clearly when any
required v2 identity field is missing. If v2 identity or metadata fields are
added without `config_schema_version: 2`, loading also fails so partial
migrations do not silently run as v1 configs. Optional v2 `quality_metadata`
and the v1/v2 `regression_thresholds` policy must be JSON objects when present.
(`recipe` is emitted by the harness and is not config-declarable.)
V2 records compare only within their exact identity cohort. A record that opens
a new cohort is marked `baseline_status: "initialized_new_cohort"` and
`comparison_status: "CALIBRATION_NEEDED"`; it remains ineligible until a
reviewed scheduled-main artifact is seeded explicitly. Legacy v1 configs still
run and are normalized for reporting, but their records skip rolling-baseline
comparison entirely (`baseline_status: "skipped_missing_identity"`, never
baseline eligible); only static thresholds gate them. Metric-specific threshold
policies are active. `QUALITY_BLOCKED` remains reserved for future variant
promotion policy.
The shipped Wan benchmark uses `benchmark_version: 3` because recipe schema 2
changed the recipe fingerprint by removing the legacy `benchmark_id` display
name. This intentionally opens a new comparison cohort: after deployment, a
reviewed scheduled-main full-suite `CALIBRATION_NEEDED` artifact must be seeded
once before rolling regression gating resumes for that exact identity. Static
thresholds remain active during calibration.
### Raw record (`results/perf_*.json`)
Written by `test_inference_performance.py`. One file per benchmark run.
@@ -179,6 +276,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
```jsonc
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"result_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 3,
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
"device": "NVIDIA L40S",
"num_gpus": 2,
@@ -196,12 +297,70 @@ Written by `test_inference_performance.py`. One file per benchmark run.
"max_dit_time_s": 10.0,
"max_vae_decode_time_s": 10.0
},
"regression_thresholds": {
"latency": {
"threshold_percent": 0.10,
"threshold_absolute": 1.0,
"gated": true
}
},
"commit": "<full sha>",
"run_source": "pr",
"branch": "feature/perf-change",
"pr_number": "1234",
"test_scope": "direct",
"build_url": "https://buildkite.example/build",
"build_id": "<buildkite-build-id>",
"job_id": "<buildkite-job-id>",
"timestamp": "2026-05-08T22:00:00+00:00",
"quality_metadata": { "quality_status": "canonical" },
"text_encoder_time_s": 2.141,
"dit_time_s": 8.437,
"vae_decode_time_s": 3.208
"vae_decode_time_s": 3.208,
"recipe": {
"recipe_schema_version": 2,
"benchmark": {
"benchmark_id": "wan-t2v-1.3b-2gpu",
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 3
},
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
"generation_kwargs": { "height": 480, "width": 832, "num_frames": 45 },
"inputs": { "prompt_count": 1, "prompt_sha256": ["<measured-prompt-sha256>"] },
"attention": { "requested_backend": "FLASH_ATTN", "resolved_backend": "FLASH_ATTN" }
},
"recipe_fingerprint": "<sha256>",
"hardware_profile": {
"device_type": "cuda",
"gpu_count": 2,
"gpus": [{ "name": "NVIDIA L40S", "memory_gb": 48, "compute_capability": "8.9" }],
"interconnect": "none_or_partial"
},
"hardware_profile_id": "hw-<sha256-prefix>",
"software_profile": {
"python": "3.12",
"pytorch": "2.12",
"cuda": "13.0",
"attention_backend": "FLASH_ATTN",
"flash_attention_4_enabled": true,
"container_image_version": "py3.12-cuda13.0.0",
"packages": {
"fastvideo_kernel": "0.3.2",
"flashinfer": "0.2.11",
"nvidia_cutlass_dsl": "4.5.0",
"triton": "3.4.1"
}
},
"software_profile_id": "sw-<sha256-prefix>",
"environment_metadata": {
"env": {
"IMAGE_VERSION": "py3.12-cuda13.0.0",
"FASTVIDEO_CONTAINER_IMAGE_REF": "ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-cuda13.0.0@sha256:<digest>"
}
},
"environment_fingerprint": "env-<sha256-prefix>"
}
```
@@ -213,6 +372,10 @@ result, used as the rolling-baseline source of truth.
```jsonc
{
"model_id": "wan-t2v-1.3b-2gpu",
"result_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 3,
"timestamp": "2026-05-08T22:00:00+00:00",
"commit_sha": "<full sha>",
"gpu_type": "NVIDIA L40S",
@@ -222,34 +385,81 @@ result, used as the rolling-baseline source of truth.
"text_encoder_time_s": 2.141,
"dit_time_s": 8.437,
"vae_decode_time_s": 3.208,
"regression_thresholds": {
"latency": {
"threshold_percent": 0.08,
"threshold_absolute": 0.5,
"gated": true
}
},
"recipe_fingerprint": "<sha256>",
"hardware_profile_id": "hw-<sha256-prefix>",
"software_profile_id": "sw-<sha256-prefix>",
"environment_fingerprint": "env-<sha256-prefix>",
"run_source": "pr",
"branch": "feature/perf-change",
"pr_number": "1234",
"test_scope": "direct",
"build_url": "https://buildkite.example/build",
"build_id": "<buildkite-build-id>",
"job_id": "<buildkite-job-id>",
"quality_metadata": { "quality_status": "canonical" },
"baseline_status": "compared",
"comparison_status": "PASS",
"comparison_status_reason": "Comparable baseline found with no gated regressions",
"baseline_eligible": false,
"success": true
}
```
### Compatibility with legacy records
Older records in the HF dataset may not have component timing fields. The
comparator ignores missing or `null` metrics when computing a median, and the
dashboard lists skipped plots for metric series that have no non-null values.
Records missing both `run_source` and `baseline_eligible` are treated as legacy
successful main/full-suite uploads and remain eligible for rolling baselines.
Older records in the HF dataset may not have `result_schema_version`,
component timing fields, or v2 identity/profile fields. Records without
`result_schema_version` are treated as v1. The comparator ignores missing or
`null` metrics when computing a median, and the dashboard lists skipped plots
for metric series that have no non-null values. Records missing both
`run_source` and `baseline_eligible` are treated as legacy successful
main/full-suite uploads and remain eligible for rolling baselines.
Current `perf_*.json` artifacts that lack the v2 comparison identity are
normalized for reporting but skip rolling-baseline comparison and are not marked
baseline eligible.
New v2 records compare only against the same `workload_id`, `variant_id`,
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
`software_profile_id` cohort, independent of the legacy `model_id` directory
and `gpu_type` display string. Historical v1 records remain readable for
reporting, but current legacy artifacts do not perform a `(model_id, gpu_type)`
rolling comparison or seed new rolling baselines.
`environment_metadata` and `environment_fingerprint` are audit data and are not
part of the comparison key.
The recipe prompt digests describe the prompts actually measured by the
benchmark run; extra configured prompts are ignored unless the benchmark runner
executes them.
Software profile package cohorts keep exact versions for relevant
attention/kernel packages, including FastVideo kernels, FlashAttention,
FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
## Environment variable reference
| Variable | Default | Used by | Purpose |
|---|---|---|---|
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py`, `test_inference_performance.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Performance pytest exit code. Measured static-threshold failures are attributed per record; otherwise a nonzero code reports an infrastructure error. |
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
| `FASTVIDEO_ATTENTION_BACKEND` | `auto` | `test_inference_performance.py` | Requested attention backend included in `software_profile_id`. |
| `FASTVIDEO_FA4` | `0` | `test_inference_performance.py` | FlashAttention-4 toggle included in `software_profile_id`. |
| `FASTVIDEO_PERFORMANCE_PROFILE_VERSION` | unset | `test_inference_performance.py` | Optional explicit software cohort/profile version included in `software_profile_id`. |
| `IMAGE_VERSION` | unset | `test_inference_performance.py` | CI container image/profile version included in `software_profile_id` when available. |
| `FASTVIDEO_CONTAINER_IMAGE_REF` | unset | `pr_test.py`, `launch_l40s_job.py`, `test_inference_performance.py` | Resolved CI container image ref or digest recorded in `environment_metadata` for audit without changing `software_profile_id`. |
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
## CI integration
@@ -260,17 +470,24 @@ point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
Buildkite artifact upload is in
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
Each performance build runs pytest first. If that fixed-threshold phase fails,
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
artifacts are not emitted. The dashboard still runs best-effort for
observability. When pytest passes, the rolling-baseline phase emits:
Each performance build runs pytest first. PR and direct runs only continue to
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
Markdown summaries and normalized JSON artifacts are not emitted. Scheduled
main runs set `PERF_UPLOAD_POLICY=always`, so they still run
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after pytest fails. Each raw
record is checked against its own static thresholds: a measured breach reports
`REGRESSION`, while unaffected records retain their rolling-baseline status. A
nonzero pytest exit with no attributable static-threshold breach reports
`INFRA_ERROR`. Failed records have `success=false` and are excluded from future
rolling baselines. The dashboard still runs best-effort for observability.
When the rolling-baseline phase runs, it emits:
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
per-benchmark row with current vs. baseline values for latency, throughput,
memory, text encoder time, DiT time, and VAE decode time.
* **Plotly dashboard** — `dashboard_<sha>_<ts>.html` showing time-series for
each metric grouped by `(model_id, gpu_type)`.
each metric grouped by comparison cohort.
* **Normalized records** — `normalized_perf_*.json`, one per benchmark.
Useful as input to the
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
@@ -279,11 +496,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
## Adding a new benchmark
1. Drop a new JSON config into
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
use v2 identity fields:
```json
{
"benchmark_id": "<unique-id>",
"config_schema_version": 2,
"workload_id": "<stable-workload-id>",
"variant_id": "<variant, e.g. 1.3b-sp2>",
"benchmark_version": 1,
"model": { "model_path": "...", "model_short_name": "..." },
"init_kwargs": { "num_gpus": 1, ... },
"generation_kwargs": { "num_frames": 45, ... },
@@ -299,17 +521,30 @@ observability. When pytest passes, the rolling-baseline phase emits:
"max_vae_decode_time_s": 10.0
},
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
},
"regression_thresholds": {
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
}
}
```
}
```
Legacy v1 configs without `config_schema_version` still load, but should not
gain v2 identity or metadata fields until they are migrated to
`config_schema_version: 2`. For v2 configs, `workload_id`, `variant_id`,
and `benchmark_version` are part of the comparison key; benchmark runs
fail if any of these identity fields are missing.
2. The pytest test auto-discovers all configs — no test code needed. CI
picks it up on the next `/test performance` run.
3. The first persisted main-branch run with no HF history initializes the
baseline (passes automatically). Subsequent runs compare against it. Local
and pull-request runs with no HF history also pass, but they do not seed the
shared baseline.
3. Legacy benchmarks are gated only by their static thresholds. Their current
records skip rolling-baseline comparison and are never baseline eligible.
V2 benchmarks with no exact comparable baseline report
`CALIBRATION_NEEDED`; the record remains visible but does not become
baseline eligible until a comparable scheduled-main full-suite run is
reviewed and seeded through the prepare, review, and conditional-upload
steps in the `reseed-performance-baseline` skill.
4. If the benchmark targets a GPU not currently in `thresholds`, either add
that GPU as a key or rely on the `default` block. Note that `default` is
@@ -320,10 +555,21 @@ observability. When pytest passes, the rolling-baseline phase emits:
a useful fixed gate. The rolling baseline will still track component times
when static component thresholds are omitted.
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
include only benchmark-specific deviations. Tune these independently from
the fixed thresholds when a metric is noisy or should be informational. The
fixed `thresholds` block is an absolute pytest ceiling. The
`regression_thresholds` block controls rolling-baseline comparisons against
recent scheduled-main records.
## Troubleshooting
**"No baseline for ... Initializing"** — first run for this `(model_id,
gpu_type)`. Run will pass and (if persisting) seed the first record.
**`CALIBRATION_NEEDED` / "No baseline found for exact comparable identity"** —
the v2 run passes, but its normalized record remains
`baseline_eligible=false`. Review a successful scheduled-main full-suite
normalized artifact, then follow the `reseed-performance-baseline` skill. The
utility only prepares a digest-protected manifest; the separately confirmed
upload rechecks remote state and commits the batch atomically.
**Persistent failure right after a torch / kernel / image upgrade** —
genuine regression *or* baseline drift. Compare the failing normalized record
@@ -336,5 +582,5 @@ pipelines that did not report a mapped component stage.
**Component timing is `null`** — the generated result did not include a mapped
stage in `logging_info.stages`. Check that the pipeline emits stage logging
and that the stage name is listed in `STAGE_METRIC_MAP` in
`test_inference_performance.py`.
and that the stage emits `component_metric` or is covered by the legacy
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
+24
View File
@@ -180,6 +180,30 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
--device-folder L40S_reference_videos
```
### SSIM Bootstrap Mode
Normal SSIM runs are strict: if a reference video or latent is missing, the
test fails. For new-model PRs, CI can run SSIM in bootstrap mode so missing
references are uploaded as draft artifacts for review instead of immediately
blocking on a missing canonical reference.
Buildkite enables SSIM bootstrap mode when either condition is true:
- the PR title or Buildkite message contains `[new-model]`;
- `FASTVIDEO_SSIM_BOOTSTRAP_MODE=1` is set for the Buildkite job.
Bootstrap mode passes `--ssim-bootstrap-mode` to pytest. When a generated
artifact is available, the test uploads it under the `drafts/...` namespace in
the SSIM reference repo and marks that case as expected-failed. After reviewing
the draft, promote it into the canonical reference layout:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
--quality-tier default \
--device-folder L40S_reference_videos \
--model-id <model_id>
```
## CI Integration
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
@@ -191,6 +191,9 @@ surfaces:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
color_correction_strength:
sources:
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
default_camera_rotation:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
@@ -455,6 +458,8 @@ surfaces:
guidance_scale: request.sampling.guidance_scale
guidance_scale_2: request.sampling.guidance_scale_2
guidance_rescale: request.sampling.guidance_rescale
use_embedded_guidance: request.sampling.use_embedded_guidance
true_cfg_scale: request.sampling.true_cfg_scale
boundary_ratio: request.sampling.boundary_ratio
sigmas: request.sampling.sigmas
enable_teacache: request.runtime.enable_teacache
+116
View File
@@ -0,0 +1,116 @@
# 🌊 AnyFlow Any-Step Video Distillation
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
```
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
```
for any ``t > r``.
## 📊 Model Overview
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
## ⚙️ Inference
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
## 🧠 Algorithm
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
### Stage 1 — Pretrain (flow-map central-difference)
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
- the remainder is free.
The student forward at ``(t, r)`` is trained against the central-difference target
```
target = (eps - x_0) - (t - r) · dF/dt
```
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
### Stage 2 — On-policy DMD
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
## 🚀 Training Scripts
### Stage 1 — pretrain
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
```
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
- Global batch size: 32 (8 GPUs × 4 per-GPU)
- Learning rate: 5e-5
- Flow shift: 5.0
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
- ``weight_type``: ``beta08``
- ``fuse_guidance_scale``: 3.0
- Training steps: 6000
### Stage 2 — on-policy
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
```
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
**Key configuration**:
- Global batch size: 8 (8 GPUs × 1 per-GPU)
- Learning rate: 2e-6
- Flow shift: 5.0
- ``student_sample_steps``: 4
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
- ``real_score_guidance_scale``: 3.0
- ``generator_update_interval``: 5 (DMD2 alternation)
- Training steps: 4000
## 🔌 Loading published AnyFlow checkpoints
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
## 🧭 Note on ``fuse_guidance_scale``
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
```
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
```
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
+17
View File
@@ -74,6 +74,23 @@ uv pip install ninja
python setup.py install
```
### Flash Attention 4 (opt-in)
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
runtime on some GPU/shape combinations. To use FA4, install the pinned
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
```bash
export FASTVIDEO_FA4=1
```
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
cannot serve there: grad-enabled (training) attention (FA4's backward requires
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
FastVideo fails loudly instead of silently falling back.
### FP4 Flash Attention 4 (Blackwell only)
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
+2
View File
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
+91
View File
@@ -0,0 +1,91 @@
# Training Trackers
FastVideo can send training metrics and validation media to Weights & Biases
or SwanLab. Tracking runs only on global rank 0, and local tracker files are
stored under `<output_dir>/tracker`.
## Supported Trackers
| Value | Backend | Installation |
|-------|---------|--------------|
| `wandb` | Weights & Biases | Included with FastVideo |
| `swanlab` | SwanLab | Install the optional `swanlab` dependency |
| `none` | Disable external tracking | No additional package |
You can enable more than one backend, for example `trackers: [wandb, swanlab]`.
Metrics and validation media are converted to the artifact type required by
each backend.
## Install SwanLab
For a published FastVideo installation, install the SwanLab extra:
```bash
uv pip install "fastvideo[swanlab]"
```
For an editable source checkout, include the same extra during installation:
```bash
uv pip install -e ".[swanlab]"
```
If FastVideo is already installed, you can install the compatible SDK directly:
```bash
uv pip install "swanlab>=0.6.7"
```
Authenticate once before starting a training run:
```bash
swanlab login
```
See the [SwanLab login documentation](https://docs.swanlab.cn/en/api/cli-swanlab-login.html)
for non-interactive and self-hosted setups.
## Configure Tracking
Select SwanLab in the YAML config used by the modular training framework:
```yaml
training:
checkpoint:
output_dir: outputs/my_run
tracker:
trackers: [swanlab]
project_name: my_project
run_name: my_run
```
To log to both supported services:
```yaml
training:
tracker:
trackers: [wandb, swanlab]
project_name: my_project
run_name: my_run
```
An empty or omitted `trackers` list selects W&B when `project_name` is set.
Use an explicit `none` entry to disable external tracking:
```yaml
training:
tracker:
trackers: [none]
```
## Validation Videos
SwanLab currently accepts GIF video artifacts. FastVideo converts validation
MP4 files and in-memory video arrays to GIF automatically before logging them.
For video files, FastVideo uses the sampling frame rate supplied by the caller,
or the source file's frame rate when no value is supplied. In-memory arrays use
the frame rate supplied by the caller. Both forms fall back to 16 FPS when no
frame rate is available.
For details about configuring validation callbacks, see
[Training Infrastructure](train_infra.md#callbacks-pluggable-hooks).
+49
View File
@@ -161,6 +161,21 @@ training:
decay_interval_steps: 0
```
`training.data.data_path` can also mix multiple preprocessed datasets by using a mapping from dataset path to repeat count:
```yaml
training:
data:
data_path:
data/zeldam2-clean: 1
data/multi3d_games: 2
```
The repeat count duplicates that dataset's parquet file list before shuffling/sampling, so the example above trains with roughly twice as much `multi3d_games` exposure as `zeldam2-clean`. Paths are just suggested locations; use any local path that contains a FastVideo preprocessed parquet dataset.
See [Training Trackers](trackers.md) to configure Weights & Biases or SwanLab,
including SwanLab installation and authentication.
### `callbacks` — Pluggable hooks
Callbacks run at specific points in the training loop (before/after optimizer
@@ -323,6 +338,40 @@ Self-Forcing inherits all DMD2 parameters, plus:
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
| `start_gradient_frame` | `0` | Frame index where gradients begin |
### Streaming Long Tuning
`StreamingLongTuningMethod` extends Self-Forcing for LongLive-style rollouts. It
keeps a streaming state, generates overlapping chunks, and trains only the new
frames while preserving context from earlier chunks.
For the MatrixGame2/Zelda world-model example, self-forcing and long tuning are
separate runs: first train or load the 1k-step self-forcing checkpoint using
`examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml`,
then run
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
from that checkpoint for the 3k-step streaming long-tuning stage.
```yaml
method:
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
streaming_reencode_overlap_anchor: true
streaming_anchor_inject_k: 1
streaming_require_full_blocks: true
multi_phased_distill_schedule:
- stage: streaming_long
start_step: 0
end_step: 3000
num_latent_t: 39
streaming_training: true
```
See
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
for a complete MatrixGame2/Zelda configuration.
---
## Callbacks
@@ -0,0 +1,64 @@
import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
def _env_int(name: str, default: int) -> int:
return int(os.getenv(name, str(default)))
def _env_float(name: str, default: float) -> float:
return float(os.getenv(name, str(default)))
def main():
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
generator = VideoGenerator.from_pretrained(
model_name,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
override_pipeline_cls_name="DreamXWorldPipeline",
)
prompt = os.getenv(
"DREAMX_WORLD_PROMPT",
"A cinematic first-person drive through a futuristic coastal city at "
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
)
image_path = os.getenv(
"DREAMX_WORLD_IMAGE_PATH",
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
)
kwargs = {
"output_path": OUTPUT_PATH,
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
"action_speed_list": [
float(value)
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
],
}
if image_path:
kwargs["image_path"] = image_path
try:
generator.generate_video(prompt, **kwargs)
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+140
View File
@@ -0,0 +1,140 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import contextlib
import os
import re
DEFAULT_PROMPTS = [
"a photo of a cat",
(
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"
),
]
def _safe_filename(text: str, max_len: int = 100) -> str:
"""Make a stable, filesystem-friendly filename base."""
s = text[:max_len].strip()
s = s.replace(os.sep, "_")
if os.altsep:
s = s.replace(os.altsep, "_")
s = re.sub(r"\s+", " ", s)
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
s = s.strip(" .")
return s or "prompt"
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
if not os.path.isdir(out_dir):
return
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
with contextlib.suppress(FileNotFoundError):
os.remove(os.path.join(out_dir, fn))
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
)
p.add_argument(
"--model-path",
default="official_weights/FLUX.1-dev",
help="Local Diffusers checkpoint dir or HF repo id.",
)
p.add_argument(
"--out-dir",
"--outdir",
default="outputs/flux_dev/samples",
help="Directory for saved PNG outputs.",
)
p.add_argument(
"--prompt",
action="append",
default=None,
help="Prompt. Repeat for multiple images.",
)
p.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
)
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
p.add_argument("--height", type=int, default=1024, help="Output height.")
p.add_argument("--width", type=int, default=1024, help="Output width.")
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
return p.parse_args()
def main() -> None:
args = parse_args()
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
from fastvideo import VideoGenerator
os.makedirs(args.out_dir, exist_ok=True)
init_kwargs = {
"num_gpus": args.num_gpus,
"workload_type": "t2i",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
"use_fsdp_inference": False,
}
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
**init_kwargs,
)
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = (
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
)
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
generation_kwargs = {
"output_path": output_path,
"height": args.height,
"width": args.width,
"num_frames": 1,
"fps": 1,
"num_inference_steps": args.steps,
"guidance_scale": args.guidance,
"use_embedded_guidance": True,
"true_cfg_scale": 1.0,
"seed": seed,
"save_video": True,
}
generator.generate_video(prompt, **generation_kwargs)
print(f"[flux] done. outputs written to: {args.out_dir}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+107
View File
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image text-to-image generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
text-to-image generation command, saved as a PNG."
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--output",
default="image_output/landscape.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default=("A beautiful landscape photography with rolling hills, "
"a winding river, and a vibrant sunset in the background. "
"Warm golden light, photorealistic style."),
help="Text prompt.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
IMAGE_PATH = "assets/girl.png"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A woman stands up and walks away"
)
_ = generator.generate_video(
prompt,
image_path=IMAGE_PATH,
output_path=OUTPUT_PATH,
save_video=True,
height=1024,
width=1024,
num_frames=121,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
if __name__ == "__main__":
main()
+120
View File
@@ -0,0 +1,120 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
want a minimal edit command (text + image -> edited image), saved as a PNG."
GLM-Image is a single unified pipeline: passing a condition image switches it
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
write pass), so the generator config is identical to `basic_glm_image.py` — the
`inputs.pil_image` on the request is what selects the edit mode.
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--image",
default="assets/images/couple.jpg",
help="Condition image to edit.",
)
parser.add_argument(
"--output",
default="image_output/edited.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default="Change the background to a snowy mountain landscape at golden hour.",
help="Edit instruction.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
condition = Image.open(args.image).convert("RGB")
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
# The pipeline is registered as t2i; passing inputs.pil_image below switches
# it to the edit path.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
inputs=InputConfig(pil_image=condition),
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+69
View File
@@ -0,0 +1,69 @@
# LTX-2.3 distilled inference configs
Ready-to-run `fastvideo generate` run configs for the LTX-2.3
distilled model (`FastVideo/LTX-2.3-Distilled-Diffusers`), covering both
workloads (t2v / i2v), both two-stage step schedules (`5+2`, `8+3` = denoise
+ refine), and four resolutions.
```bash
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
```
Each config is self-contained (no preset registry needed): the two-stage
refine is wired via `generator.pipeline.preset_overrides.refine`, and the
base sampling knobs live under `request.sampling`. The refine upsampler
auto-resolves from the model's `spatial_upscaler`.
## Configs
| workload | schedule | resolution (HxW) | file |
|---|---|---|---|
| t2v | 5+2 | 1280x832 | `t2v_5s2_1280x832.yaml` |
| t2v | 5+2 | 1024x1536 | `t2v_5s2_1024x1536.yaml` |
| t2v | 5+2 | 768x1280 | `t2v_5s2_768x1280.yaml` |
| t2v | 5+2 | 512x768 | `t2v_5s2_512x768.yaml` |
| t2v | 8+3 | 1280x832 | `t2v_8s3_1280x832.yaml` |
| t2v | 8+3 | 1024x1536 | `t2v_8s3_1024x1536.yaml` |
| t2v | 8+3 | 768x1280 | `t2v_8s3_768x1280.yaml` |
| t2v | 8+3 | 512x768 | `t2v_8s3_512x768.yaml` |
| i2v | 5+2 | 1280x832 | `i2v_5s2_1280x832.yaml` |
| i2v | 5+2 | 1024x1536 | `i2v_5s2_1024x1536.yaml` |
| i2v | 5+2 | 768x1280 | `i2v_5s2_768x1280.yaml` |
| i2v | 5+2 | 512x768 | `i2v_5s2_512x768.yaml` |
| i2v | 8+3 | 1280x832 | `i2v_8s3_1280x832.yaml` |
| i2v | 8+3 | 1024x1536 | `i2v_8s3_1024x1536.yaml` |
| i2v | 8+3 | 768x1280 | `i2v_8s3_768x1280.yaml` |
| i2v | 8+3 | 512x768 | `i2v_8s3_512x768.yaml` |
## Overriding without editing a file
Dotted overrides (prefixes `generator.` / `request.`) let you tweak any field:
```bash
# swap prompt
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml \
--request.prompt "a red fox running through fresh snow"
# change output path / gpu count
fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml \
--request.output.output_path outputs/preview.mp4 \
--generator.engine.num_gpus 4
```
## i2v
The `i2v_*` configs take a first-frame image via
`request.extensions.ltx2_images` (`[[path, frame_offset, weight]]`). Edit the
path in the file, or override it:
```bash
fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml \
--request.extensions.ltx2_images '[["/data/portrait.jpg", 0, 1.0]]'
```
## Schedules
`5+2` is the fast preview schedule; `8+3` is the higher-quality distilled
recipe. Refine (`preset_overrides.refine.num_inference_steps`) only accepts 2
or 3 steps.
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1024x1536.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_1024x1536.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1280x832.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_1280x832.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_512x768.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_512x768.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_768x1280.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_768x1280.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1024x1536.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_1024x1536.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_1280x832.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_512x768.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_512x768.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_768x1280.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_768x1280.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1024x1536.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_1024x1536.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1280x832.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_1280x832.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_512x768.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_768x1280.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_768x1280.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1024x1536.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_1024x1536.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_1280x832.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_512x768.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_512x768.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_768x1280.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_768x1280.mp4
save_video: true
@@ -0,0 +1,78 @@
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
#
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
# the student matches an EMA copy of itself at the next timestep, all under
# clean-history teacher forcing.
#
# All three roles initialize from the SAME checkpoint (the teacher-forcing
# AR-diffusion model). Point init_from at that checkpoint for a real run.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_cd
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_cd_shift5
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
pipeline:
flow_shift: 5
@@ -0,0 +1,105 @@
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
#
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
# rollout from pure noise. Teacher provides the real score, critic
# learns the fake score; both inherited from DMD2Method.
#
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
# or with the NVIDIA-released checkpoint
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
# the paper weights (the delta_embedder rename is handled by the
# param_names_mapping in WanVideoArchConfig).
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: <PATH_TO_PRETRAIN_CKPT>
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 3.0
dmd_denoising_steps: [999, 937, 833, 624]
warp_denoising_step: false
# AnyFlow rollout knobs.
student_sample_steps: 4
use_mean_velocity: true
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
# Critic optimizer (DMD2 inherited).
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
attn_kind: vsa
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_onpolicy
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_onpolicy
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
@@ -0,0 +1,83 @@
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
#
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
# can be sampled at arbitrary NFE in the on-policy stage.
#
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
diffusion_ratio: 0.5
consistency_ratio: 0.25
epsilon: 5 # finite-difference step in absolute train-timestep units
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
fuse_guidance_scale: 3.0
# shift is taken from pipeline.flow_shift below.
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 4
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 6000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_pretrain
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_pretrain
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
# Enable AnyFlow dual-timestep conditioning. The student loads from
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
# get initialized identically to time_embedder via deep-copy in
# WanTimeTextImageEmbedding.__init__.
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
@@ -100,9 +100,6 @@ callbacks:
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 81
# Validation/inference uses standard CFG in both clean and Self-Forcing,
# so this directly matches Self-Forcing guidance_scale=3.0.
guidance_scale: 3.0
pipeline:
flow_shift: 5
@@ -0,0 +1,73 @@
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model with a block size of 1 frame
# - Training: each frame gets its own independent noise level (frame-wise
# diffusion forcing), versus the chunk-wise variant that shares one noise
# level across num_frames_per_block frames.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 1
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_dfsft_framewise
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
@@ -0,0 +1,72 @@
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
# denoises the current block while attending to *clean* history (clean_x),
# not its own noisy rollout.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_tfsft
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
+96 -9
View File
@@ -1,32 +1,119 @@
# World-Model: Matrix-Game 2.0 I2V
Three training scenarios for the Matrix-Game 2.0 I2V world model on the
new YAML-driven trainer (`fastvideo/train/entrypoint/train.py`).
Training scenarios for the Matrix-Game 2.0 I2V world model on Solaris (Minecraft)
data and Zelda data, using the new YAML-driven trainer
(`fastvideo/train/entrypoint/train.py`).
## Solaris Configs
| Config | Method | Student | Notes |
|---|---|---|---|
| `finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
| `dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
| `self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
| `solaris/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
| `solaris/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
| `solaris/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Matrix-Game 2.0 DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
## Zelda Configs
| Config | Method | Student | Notes |
|---|---|---|---|
| `zelda/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Zelda bidirectional I2V finetuning from `FastVideo/Matrix-Game-2.0-Base-Diffusers`. Uses 33-frame clips and Zelda validation with action overlays. |
| `zelda/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Zelda causal Diffusion-Forcing SFT from `mignonjia/mg_bidirectional_zelda`. Uses the same Zelda data, resolution, optimizer, and validation defaults as the Zelda finetune config. |
| `zelda/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Zelda DMD/Self-Forcing distillation; student init = `mignonjia/mg_causal_zelda`, teacher = bidirectional (`mignonjia/mg_bidirectional_zelda`), critic = bidirectional. |
| `zelda/streaming_long_tuning_causal_i2v.yaml` | `StreamingLongTuningMethod` | `MatrixGame2CausalModel` | LongLive-style streaming long tuning from the 1k-step Zelda self-forcing checkpoint. |
Zelda world-model distillation is a two-run workflow: first run
`zelda/self_forcing_causal_i2v.yaml` to train or load the 1k-step
self-forcing checkpoint (`mignonjia/mg_sf_distilled_zelda_1k_steps`), then run
`zelda/streaming_long_tuning_causal_i2v.yaml` for the 3k-step streaming
long-tuning stage. The long-tuning YAML starts from that 1k-step checkpoint; it
does not run the short self-forcing stage inside the same config.
## Zelda Training Data
The Zelda training configs use `data/zeldam2-clean` as a suggested local path.
Download the dataset from Hugging Face before running those configs:
```bash
python scripts/huggingface/download_hf.py \
--repo_id mignonjia/zeldam2-clean \
--local_dir data/zeldam2-clean \
--repo_type dataset
```
You can store the dataset elsewhere; update `training.data.data_path` in the
YAML to point at that location.
## Multi3D Training Data
`zelda/finetune_i2v.yaml` includes an optional, commented-out Multi3D entry.
Enable it only when you want to mix Zelda with multi-game data from
`data/multi3d_games`. You can store this dataset anywhere; before enabling it,
update the matching commented `training.data.data_path` key in the YAML to the
correct location.
To mix datasets in a training YAML, set `training.data.data_path` to a
path-to-repeat-count mapping. For example, `zelda/finetune_i2v.yaml` can use
`data/zeldam2-clean: 1` and `# data/multi3d_games: 10`; uncommenting the
Multi3D entry repeats the multi-game parquet list ten times before training
samples are shuffled.
## World Model Validation Data
The Zelda validation configs expect a small public validation bundle under
`data/zelda_validation_data`.
Download it from Hugging Face before running the Zelda scenarios:
```bash
python scripts/huggingface/download_hf.py \
--repo_id mignonjia/zelda_validation_data \
--local_dir data/zelda_validation_data \
--repo_type dataset
```
The bundle contains `validation_zelda.json`, `images/`, and `actions/`.
The Zelda configs point
`callbacks.validation.dataset_file` at
`data/zelda_validation_data/validation_zelda.json`.
## Usage
### Solaris
```bash
bash examples/train/run.sh \
examples/train/scenario/worldmodel/finetune_i2v.yaml
examples/train/scenario/worldmodel/solaris/finetune_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/self_forcing_causal_i2v.yaml
examples/train/scenario/worldmodel/solaris/self_forcing_causal_i2v.yaml
```
### Zelda
```bash
# Finetuning / DFSFT
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/finetune_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/dfsft_causal_i2v.yaml
# Distillation / long tuning
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml
```
Override any field on the command line:
```bash
bash examples/train/run.sh \
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml \
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml \
--training.distributed.num_gpus 8 \
--training.optimizer.learning_rate 1e-5
```
@@ -97,4 +97,4 @@ callbacks:
guidance_scale: 6.0
pipeline:
flow_shift: 5
flow_shift: 5
@@ -0,0 +1,94 @@
# Diffusion-Forcing SFT: Zelda world model I2V Causal
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path:
data/zeldam2-clean: 1
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 2.0e-5
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 60000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_causal_dfsft
training_state_checkpointing_steps: 5000
checkpoints_total_limit: 3
tracker:
entity: hapo-exp
project_name: mg_1.3b_zelda
run_name: zelda_causal_dfsft
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 200
sampling_steps: [40]
sampling_timesteps: [1000, 975, 950, 925, 900, 875, 850, 825, 800, 775,
750, 725, 700, 675, 650, 625, 600, 575, 550, 525,
500, 475, 450, 425, 400, 375, 350, 325, 300, 275,
250, 225, 200, 175, 150, 125, 100, 75, 50, 25]
num_frames: 33
overlay_actions: true
guidance_scale: 6.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
@@ -0,0 +1,88 @@
# Matrix-Game 2.0 Zelda + multi-game I2V finetune.
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: FastVideo/Matrix-Game-2.0-Base-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path:
data/zeldam2-clean: 1
# data/multi3d_games: 10
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.0 # unused for MatrixGame2 I2V; no text_embedding CFG dropout
seed: 42
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 2.0e-5
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 60000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_with_mg_init
training_state_checkpointing_steps: 5000
checkpoints_total_limit: 3
tracker:
project_name: mg_1.3b_zelda
run_name: zelda_with_mg_init
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_i2v_pipeline.MatrixGame2I2VPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 200
sampling_steps: [40]
num_frames: 33
overlay_actions: true
guidance_scale: 6.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
@@ -0,0 +1,120 @@
# Self-forcing distillation: Zelda world model I2V Causal
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_causal_zelda
trainable: true
teacher:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
chunk_size: 3
student_sample_type: sde
same_step_across_blocks: true
last_step_only: false
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
# Critic optimizer
fake_score_learning_rate: 3.0e-7
fake_score_betas: [0.9, 0.95]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/zeldam2-clean
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1001
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 3.0e-6
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/zelda_causal_self_forcing
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
project_name: wangame_sf
run_name: mg2_self_forcing_9_latents
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 100
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 153
overlay_actions: true
keyboard_value_scale: 1.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
@@ -0,0 +1,140 @@
# MatrixGame2 I2V LongLive-style streaming distillation: Zelda world model I2V Causal
# Student init from self forcing after 1k steps
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_sf_distilled_zelda_1k_steps
trainable: true
# transformer_override_safetensor: outputs/matrixgame_dmd/checkpoints/mg_zelda_sf_m2/checkpoint-1000_weight_only/ema/generator_ema.safetensors
teacher:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
rollout_mode: simulate
generator_update_interval: 5
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
chunk_size: 3
student_sample_type: sde
same_step_across_blocks: true
last_step_only: false
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
streaming_training: true
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
streaming_reencode_overlap_anchor: true
streaming_anchor_inject_k: 1
streaming_require_full_blocks: true
multi_phased_distill_schedule:
- stage: streaming_long
start_step: 0
end_step: 3000
num_latent_t: 39
streaming_training: true
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
# Critic optimizer
fake_score_learning_rate: 3.0e-7
fake_score_betas: [0.9, 0.95]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/zeldam2-clean
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1001
num_latent_t: 39
num_height: 480
num_width: 832
num_frames: 153
optimizer:
learning_rate: 3.0e-6
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/zelda_causal_long_tuning
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
project_name: wangame_sf
run_name: mg2_39only_streaming_long
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 100
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 153
overlay_actions: true
keyboard_value_scale: 1.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
+31 -5
View File
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
endif()
if(NOT GPU_BACKEND)
set(GPU_BACKEND "CUDA")
endif()
if(GPU_BACKEND STREQUAL "ROCM")
enable_language(HIP)
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
if(_FASTVIDEO_USER_CUDA_ARCH)
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
# Only numeric spellings translate; keywords like "native"/"all" would
# otherwise be mangled into nonsense ("nativ.e").
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
message(FATAL_ERROR
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
"supported. Use a numeric arch (e.g. 90a, 121), set "
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
"auto-detect from the visible GPU.")
endif()
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
@@ -173,6 +185,14 @@ else()
endif()
endif()
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
# there, and tk's base_types.cuh brace-initializes signed-char vector members
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
set(ENABLE_TK_KERNELS OFF)
endif()
if(ENABLE_TK_KERNELS)
message(STATUS "ThunderKittens kernels: ENABLED")
else()
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
"--expt-relaxed-constexpr"
"-Xcompiler=-fno-strict-aliasing"
"-Xcompiler=-fPIC"
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
# is a no-op on x86_64, where char is already signed.
"-Xcompiler=-fsigned-char"
"-DTORCH_COMPILE"
"-Xnvlink=--verbose"
"-Xptxas=--verbose"
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
endif()
# One-look answer to "what is this build producing?" — kept last so it is the
# final thing configure prints. The per-kernel matrix lives in README.md.
message(STATUS "============== fastvideo-kernel build summary ==============")
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
message(STATUS "============================================================")
+36
View File
@@ -2,6 +2,42 @@
CUDA kernels for FastVideo video generation.
## Kernel inventory
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
| Extension | Kernels | Sources | GPU arch | Build gate |
|---|---|---|---|---|
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|---|---|---|---|---|---|
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
Notes:
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
## Installation
### Standard Installation (Local Development)
+17
View File
@@ -46,6 +46,23 @@ fi
if git rev-parse --git-dir >/dev/null 2>&1; then
git submodule update --init --recursive include/cutlass include/tk
fi
# Fail fast with a clear message if the headers are still missing (e.g. a
# Docker context that excluded .git AND the submodule contents) instead of
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
if [ ! -d include/cutlass/include ]; then
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
echo " CUTLASS sources must already be present (run" >&2
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
echo " in the source checkout, or include them in the build context)." >&2
exit 1
fi
if [ ! -d include/tk/include ]; then
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
fi
# Install build dependencies
uv pip install scikit-build-core cmake ninja
+27
View File
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.api.sampling_param import SamplingParam
@dataclass
class FluxSamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 1024
width: int = 1024
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 3.5
use_embedded_guidance: bool = True
true_cfg_scale: float = 1.0
+16
View File
@@ -90,6 +90,10 @@ class SamplingParam:
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
use_embedded_guidance: bool = False
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
true_cfg_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
@@ -325,6 +329,18 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--use-embedded-guidance",
action="store_true",
default=SamplingParam.use_embedded_guidance,
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
)
parser.add_argument(
"--true-cfg-scale",
type=float,
default=SamplingParam.true_cfg_scale,
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
)
parser.add_argument(
"--boundary-ratio",
type=float,

Some files were not shown because too many files have changed in this diff Show More