Compare commits

...
21 Commits
Author SHA1 Message Date
SolitaryThinker f0e727f5ee [misc]: point DreamX-World at published FastVideo Diffusers repos
Advertise FastVideo/DreamX-World-5B-Cam-Diffusers and
FastVideo/DreamX-World-5B-Diffusers as the loadable ids in the registry
(raw GD-ML path detection is kept via the pattern detectors), default the
example and local tests to the hub ids so they no longer skip when local
converted dirs are absent, and drop the now-redundant pipeline class
override in the example (model_index.json carries DreamXWorldPipeline).
2026-07-06 12:35:40 -07:00
SolitaryThinker de6a60ac90 [test]: randomize the zero-init head so DreamX AR tiny parity is not vacuous 2026-07-05 14:39:39 -07:00
SolitaryThinker 1b4fdc2c41 [refactor]: DreamX review wave 3 — AR DiT onto FastVideo layer primitives + real param_names_mapping 2026-07-05 14:39:39 -07:00
SolitaryThinker 385cd65abb [bugfix]: DreamX review wave 2 — AR pipeline actually encodes its conditioning image 2026-07-05 14:39:39 -07:00
SolitaryThinker 156d86b70c [bugfix]: DreamX review wave 1 — detector exclusivity, converter self-heal, SP hard-fail guard 2026-07-05 14:39:39 -07:00
Suckl bab79fb5f6 Add DreamX World 5B AR pipeline 2026-07-05 14:39:39 -07:00
Suckl 073a9e78f2 Add DreamX World 5B Cam pipeline 2026-07-05 14:39:39 -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
87 changed files with 7063 additions and 369 deletions
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `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:
@@ -148,8 +148,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 +182,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,7 +191,7 @@ 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`:
@@ -200,8 +199,8 @@ Then back up only 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
)
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
Load the last 5 successful records for the target:
```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",
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
Use the shared storage helper so the path and repo type match CI:
```python
from hf_store import upload_record
from fastvideo.performance.hf_store import upload_record
upload_record("<local_record_path>", record, strict=True)
```
@@ -460,7 +459,7 @@ 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,
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
`sanitize()`, and `upload_record()`.
- `fastvideo/tests/performance/test_inference_performance.py` — source result
JSON schema.
@@ -1,5 +1,9 @@
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
"description": "Wan2.1 T2V 1.3B inference performance",
"model": {
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
+26
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"
+19 -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"
+2 -1
View File
@@ -125,7 +125,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 +136,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
+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
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
+8 -3
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
@@ -440,6 +440,8 @@ export default function App() {
<th>Throughput</th>
<th>Memory</th>
<th>Worst</th>
<th>Exceeded</th>
<th>Failing</th>
</tr>
</thead>
<tbody>
@@ -465,6 +467,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>
@@ -2,6 +2,12 @@ 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;
@@ -15,7 +21,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;
+4 -2
View File
@@ -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" && \
+2
View File
@@ -103,6 +103,7 @@ can merge a PR.
|---|---|---|
| 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 +145,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` |
+126 -29
View File
@@ -72,7 +72,10 @@ 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
@@ -92,25 +95,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
@@ -156,9 +162,22 @@ headroom and almost never need touching.
`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 each available metric, and evaluates the current run with the metric's
rolling regression policy. 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.
@@ -172,6 +191,45 @@ agent skill to advance the rolling median.
## 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-1.3b",
"variant_id": "canonical",
"benchmark_version": 1
}
```
`benchmark_id` is still required in this phase because raw artifact names,
generated-video directories, normalized record paths, and the current rolling
baseline comparator still depend on it. The v2 identity fields are config
metadata that make the measured workload explicit:
| Field | Purpose |
|---|---|
| `workload_id` | Stable benchmark family, such as `wan-t2v-1.3b`. |
| `variant_id` | Intentional recipe family, such as `canonical`. |
| `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 metadata fields
reserved for follow-up work, such as `recipe`, `metric_threshold_policy`, and
`quality_metadata`, must be JSON objects when present.
Recipe fingerprinting, hardware/software profile IDs, exact-identity
comparison, metric-specific threshold policy behavior, promoted baselines, and
dashboard regrouping are separate follow-up changes. Until those land, rolling
baseline comparison remains keyed by `(model_id, gpu_type)`.
### Raw record (`results/perf_*.json`)
Written by `test_inference_performance.py`. One file per benchmark run.
@@ -179,6 +237,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
```jsonc
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
"device": "NVIDIA L40S",
"num_gpus": 2,
@@ -196,6 +258,13 @@ 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>",
"pr_number": "1234",
"timestamp": "2026-05-08T22:00:00+00:00",
@@ -222,6 +291,13 @@ 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
}
},
"success": true
}
```
@@ -238,18 +314,17 @@ successful main/full-suite uploads and remain eligible for rolling baselines.
| 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. |
| `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` | 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`. |
| `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_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
## CI integration
@@ -261,9 +336,11 @@ 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:
PR/direct runs skip `compare_baseline.py` because they only upload passing
records. Scheduled-main runs still execute `compare_baseline.py` with
`PERF_PYTEST_RC` set so the failed canonical attempt is visible in normalized
JSON and dashboard history. The dashboard 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
@@ -279,11 +356,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": "canonical",
"benchmark_version": 1,
"model": { "model_path": "...", "model_short_name": "..." },
"init_kwargs": { "num_gpus": 1, ... },
"generation_kwargs": { "num_frames": 45, ... },
@@ -299,9 +381,17 @@ 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`.
2. The pytest test auto-discovers all configs — no test code needed. CI
picks it up on the next `/test performance` run.
@@ -320,6 +410,13 @@ 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,
@@ -336,5 +433,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`.
@@ -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
+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 | ❌ | ❌ | ✅ | ⭕ | ⭕ |
@@ -0,0 +1,70 @@
"""DreamX-World-5B-Cam camera-controlled video generation.
Uses the pre-converted Diffusers checkpoint FastVideo/DreamX-World-5B-Cam-Diffusers.
To convert the raw GD-ML/DreamX-World-5B-Cam checkpoint yourself, see
scripts/checkpoint_conversion/dreamx_world_to_diffusers.py.
"""
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,
)
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()
+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
+39 -23
View File
@@ -1,15 +1,39 @@
# SPDX-License-Identifier: Apache-2.0
import importlib.util
import os
import torch
import torch.nn.functional as F
from dataclasses import dataclass
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
from fastvideo import envs
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
# kernel package's FASTVIDEO_VSA_CUTEDSL: its CuTeDSL kernels JIT-compile per
# shape family and can fail at runtime on some arch/shape combinations, so it
# is never auto-selected just because it is installed. Below sm90 a capability
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
# JIT, observed on sm_89).
if envs.FASTVIDEO_FA4:
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
fa_version = "4"
except ImportError:
else:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
@@ -21,6 +45,12 @@ except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
try:
if importlib.util.find_spec("flash_attn.cute") is not None:
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
"set FASTVIDEO_FA4=1 to use it for inference.")
except ImportError:
pass
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
# already a registered torch.library custom op, so dynamo treats it as a
@@ -87,9 +117,10 @@ if fa_version in ("2", "3"):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` is already a torch.library custom op
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
# passthrough is enough — no extra registration needed.
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
# registered torch.library custom op (with an FA4 backward on sm90+;
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
# is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
@@ -99,17 +130,6 @@ else:
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
_WARNED_NON_FA_DTYPE = False
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
@@ -271,12 +291,8 @@ class FlashAttentionImpl(AttentionImpl):
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
orig_dtype = query.dtype
if orig_dtype not in (torch.float16, torch.bfloat16):
global _WARNED_NON_FA_DTYPE
if not _WARNED_NON_FA_DTYPE:
_WARNED_NON_FA_DTYPE = True
logger.warning(
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
"kernel and restoring on output.", orig_dtype)
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
f"bfloat16 for the kernel and restoring on output.")
query = query.to(torch.bfloat16)
key = key.to(torch.bfloat16)
value = value.to(torch.bfloat16)
+66 -73
View File
@@ -4,10 +4,9 @@ import functools
from collections.abc import Callable
import torch
from flash_attn import flash_attn_func as _flash_attn_2_func
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
logger = init_logger(__name__)
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
except ImportError:
# flash_attn.cute (FA4) is simply not installed -- expected on builds
# without it; callers fall back to FA3/FA2 quietly.
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
raise
except Exception as e:
# flash_attn.cute IS installed but failed to import -- almost always an
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
# ImportError). This is fixable by pinning a compatible
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
# callers fall back to FA3/FA2 instead of crashing worker init.
# callers can handle it uniformly.
logger.warning(
"flash_attn.cute (FA4) is installed but failed to import (%r); "
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
"restore FA4.", e)
"flash_attn.cute (FA4) is installed but failed to import (%r). "
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
"compatible nvidia-cutlass-dsl to restore FA4.", e)
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
else:
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
try:
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
# Optional so FA4-only installs can still import this module.
from flash_attn import flash_attn_func as _flash_attn_2_func
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
except ImportError:
_flash_attn_2_func = None
_flash_attn_2_varlen_func = None
def _check_dropout(dropout_p: float) -> None:
if dropout_p != 0.0:
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
@functools.cache
def _sm90_or_newer() -> bool:
return current_platform.has_device_capability(90)
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
if _sm90_or_newer():
return False
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
# capability gate, not a runtime fallback):
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
if q.shape[-2] != k.shape[-2]:
return True
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
if fa2_func is None:
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
"not installed.")
return fa2_func
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_forward",
mutates_args=(),
@@ -243,70 +277,6 @@ torch.library.register_autograd(
)
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
# fail MLIR op creation at runtime even though the import succeeded (observed:
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
class _FA4Policy:
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
FA4 is skipped when:
* a previous call failed at runtime -- CuTeDSL JIT compilation is
shape-dependent, so the first failure disables FA4 for the rest of
the process instead of retrying a broken JIT on every call; or
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
dies on its arch check) and is unvalidated for training in this repo
(its lse is not even allocated through our inference-shaped custom
op), so training keeps the pre-FA4 behavior: FA2 on every device.
"""
def __init__(self) -> None:
self.broken = False
def use_fa4(self, *tensors: torch.Tensor) -> bool:
if self.broken:
return False
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
def mark_broken(self, error: Exception) -> None:
if not self.broken:
self.broken = True
logger.warning(
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
"to FA2 for the rest of this process.", error)
_FA4 = _FA4Policy()
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
the first FA4 runtime failure) the call is served by ``fa2_func``.
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
failure, so it propagates without disabling FA4.
"""
def decorator(fa4_func: Callable) -> Callable:
@functools.wraps(fa4_func)
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
if _FA4.use_fa4(q, k, v):
try:
return fa4_func(q, k, v, *args, **kwargs)
except NotImplementedError:
raise
except Exception as e: # CuTeDSL compile errors surface as ValueError
_FA4.mark_broken(e)
return fa2_func(q, k, v, *args, **kwargs)
return wrapper
return decorator
@_with_fa2_fallback(_flash_attn_2_func)
def flash_attn_func(
q: torch.Tensor,
k: torch.Tensor,
@@ -317,6 +287,16 @@ def flash_attn_func(
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
if _use_fa2(q, k, v):
return _fa2_or_raise(_flash_attn_2_func)(
q,
k,
v,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
_check_dropout(dropout_p)
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
return out
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
@_with_fa2_fallback(_flash_attn_2_varlen_func)
def flash_attn_varlen_func(
q: torch.Tensor,
k: torch.Tensor,
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
if _use_fa2(q, k, v):
return _fa2_or_raise(_flash_attn_2_varlen_func)(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
_check_dropout(dropout_p)
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
q,
+23 -12
View File
@@ -21,24 +21,35 @@ from einops import rearrange
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from fastvideo import envs
def _resolve_flash_attn_varlen_func() -> Any:
try:
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
if envs.FASTVIDEO_FA4:
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
# not ImportError: importers like bsa_attn.py treat ImportError as
# "flash-attn not installed" and silently degrade to reference kernels.
try:
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
return flash_attn_varlen_func_cute
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
return flash_attn_varlen_func_interface
except ImportError:
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_interface
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_flash
return flash_attn_varlen_func_flash
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
+4 -3
View File
@@ -1,5 +1,6 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
from fastvideo.configs.models.dits.flux_2 import Flux2Config
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
@@ -13,7 +14,7 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
]
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
@dataclass
class DreamXWorldArchConfig(WanVideoArchConfig):
"""DreamX-World DiT config with camera PRoPE control fields."""
add_control_adapter: bool = True
cam_method: str | None = "prope"
attn_compress: int = 1
cam_self_attn_layers: tuple[int, ...] | None = None
@dataclass
class DreamXWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
prefix: str = "Wan"
@dataclass
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
"""DreamX-World-5B autoregressive causal DiT config."""
model_type: str = "ti2v"
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
text_dim: int = 4096
freq_dim: int = 256
attn_compress: int = 4
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
local_attn_size: int = 12
sink_size: int = 3
num_frames_per_block: int = 3
rope_cache_policy: str = "block_relativistic"
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
# already uses FastVideo's native key names and the converter copies the
# tensors verbatim, so every rule is an identity. The rules enumerate the
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
# norm2 / head.norm are affine-free and have no parameters).
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
r"^head\.head\.(.*)$": r"head.head.\1",
r"^head\.modulation$": r"head.modulation",
})
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
@dataclass
class DreamXWorldARConfig(DreamXWorldConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
prefix: str = "Wan"
+3 -1
View File
@@ -1,6 +1,7 @@
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
@@ -16,5 +17,6 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+127
View File
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B-Cam FastVideo model configuration helpers."""
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits.dreamx_world import (DreamXWorldARArchConfig, DreamXWorldARConfig,
DreamXWorldArchConfig, DreamXWorldConfig)
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.wan import LucyEditDevConfig, t5_postprocess_text
def make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig:
"""Return the DreamX-World DiT config matching DreamX-World-5B-Cam."""
return DreamXWorldConfig(arch_config=DreamXWorldArchConfig(
num_attention_heads=24,
attention_head_dim=128,
in_channels=48,
out_channels=48,
ffn_dim=14336,
num_layers=30,
cross_attn_norm=True,
qk_norm="rms_norm_across_heads",
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=None,
))
def make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig:
"""Return the DreamX-World-5B autoregressive causal DiT config."""
return DreamXWorldARConfig(arch_config=DreamXWorldARArchConfig(
model_type="ti2v",
num_attention_heads=24,
attention_head_dim=128,
in_channels=48,
out_channels=48,
ffn_dim=14336,
num_layers=30,
cross_attn_norm=True,
qk_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=4,
cam_self_attn_layers=tuple(range(30)),
local_attn_size=12,
sink_size=3,
num_frames_per_block=3,
))
def make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig:
"""Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam."""
return LucyEditDevConfig().vae_config
def make_dreamx_world_5b_cam_text_encoder_config() -> T5Config:
"""Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam."""
return T5Config(
arch_config=T5ArchConfig(
vocab_size=256384,
d_model=4096,
d_kv=64,
d_ff=10240,
num_layers=24,
num_decoder_layers=None,
num_heads=64,
relative_attention_num_buckets=32,
dropout_rate=0.0,
text_len=512,
feed_forward_proj="gelu",
is_encoder_decoder=False,
),
prefix="umt5",
)
@dataclass
class DreamXWorld5BCamPipelineConfig(PipelineConfig):
"""Pipeline config for the first-scope DreamX-World-5B-Cam mode."""
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_cam_dit_config)
vae_config: VAEConfig = field(default_factory=make_dreamx_world_5b_cam_vae_config)
text_encoder_configs: tuple[EncoderConfig,
...] = field(default_factory=lambda: (make_dreamx_world_5b_cam_text_encoder_config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (t5_postprocess_text, ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
flow_shift: float | None = 3.0
ti2v_task: bool = True
expand_timesteps: bool = True
vae_tiling: bool = False
vae_sp: bool = False
vae_precision: str = "fp32"
vae_decode_precision: str | None = "bf16"
dit_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.expand_timesteps = self.expand_timesteps
@dataclass
class DreamXWorld5BARPipelineConfig(DreamXWorld5BCamPipelineConfig):
"""Pipeline config for DreamX-World-5B autoregressive forcing."""
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_ar_dit_config)
flow_shift: float | None = 5.0
ti2v_task: bool = True
is_causal: bool = True
dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250)
warp_denoising_step: bool = True
context_noise: float = 0.1
num_frames_per_block: int = 3
color_correction_strength: float = 1.0
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.expand_timesteps = True
+10
View File
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
@@ -207,9 +208,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
# - "SAGE_ATTN": use Sage Attention
# - "SAGE_ATTN_THREE": use Sage Attention 3
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
# FASTVIDEO_FA4=1 as well (see below).
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# Use dedicated multiprocess context for workers.
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
+2 -2
View File
@@ -323,9 +323,9 @@ class CausalWanTransformerBlock(nn.Module):
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
+511
View File
@@ -0,0 +1,511 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldConfig
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.models.dits.wanvideo import (LayerNormScaleShift,
PatchEmbed,
WanTimeTextImageEmbedding,
WanTransformer3DModel,
WanTransformerBlock)
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.attention import LocalAttention
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.models.dits.base import BaseDiT
def _dreamx_invert_se3(transforms: torch.Tensor) -> torch.Tensor:
assert transforms.shape[-2:] == (4, 4)
rot_inv = transforms[..., :3, :3].transpose(-1, -2)
out = torch.zeros_like(transforms)
out[..., :3, :3] = rot_inv
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", rot_inv,
transforms[..., :3, 3])
out[..., 3, 3] = 1.0
return out.to(dtype=transforms.dtype)
def _dreamx_lift_k(intrinsics: torch.Tensor) -> torch.Tensor:
assert intrinsics.shape[-2:] == (3, 3)
out = torch.zeros(intrinsics.shape[:-2] + (4, 4),
device=intrinsics.device,
dtype=intrinsics.dtype)
out[..., :3, :3] = intrinsics
out[..., 3, 3] = 1.0
return out
def _dreamx_invert_k(intrinsics: torch.Tensor) -> torch.Tensor:
assert intrinsics.shape[-2:] == (3, 3)
out = torch.zeros_like(intrinsics)
out[..., 0, 0] = 1.0 / intrinsics[..., 0, 0]
out[..., 1, 1] = 1.0 / intrinsics[..., 1, 1]
out[..., 0, 2] = -intrinsics[..., 0, 2] / intrinsics[..., 0, 0]
out[..., 1, 2] = -intrinsics[..., 1, 2] / intrinsics[..., 1, 1]
out[..., 2, 2] = 1.0
return out.to(dtype=intrinsics.dtype)
def _dreamx_apply_tiled_projmat(feats: torch.Tensor,
matrix: torch.Tensor) -> torch.Tensor:
batch, num_heads, seq_len, feat_dim = feats.shape
proj_dim = matrix.shape[-1]
assert feat_dim % proj_dim == 0
if matrix.shape[1] == seq_len:
feats = feats.view(batch, num_heads, seq_len, feat_dim // proj_dim,
proj_dim)
out = torch.einsum("btij,bntpj->bntpi", matrix, feats)
return out.reshape(batch, num_heads, seq_len, feat_dim)
cameras = matrix.shape[1]
assert seq_len > cameras and seq_len % cameras == 0
feats = feats.reshape(batch, num_heads, cameras, -1,
feat_dim // proj_dim, proj_dim)
out = torch.einsum("bcij,bncpkj->bncpki", matrix, feats)
return out.reshape(batch, num_heads, seq_len, feat_dim)
def _dreamx_prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
viewmats: torch.Tensor, intrinsics: torch.Tensor):
batch, num_heads, seq_len, head_dim = q.shape
cameras = viewmats.shape[1]
assert q.shape == k.shape == v.shape
assert viewmats.shape == (batch, cameras, 4, 4)
assert intrinsics.shape == (batch, cameras, 3, 3)
assert head_dim % 4 == 0
intrinsics_norm = torch.zeros_like(intrinsics)
intrinsics_norm[..., 0, 0] = intrinsics[..., 0, 0]
intrinsics_norm[..., 1, 1] = intrinsics[..., 1, 1]
intrinsics_norm[..., 2, 2] = 1.0
proj = torch.einsum("...ij,...jk->...ik",
_dreamx_lift_k(intrinsics_norm), viewmats)
proj_t = proj.transpose(-1, -2).to(dtype=viewmats.dtype)
proj_inv = torch.einsum(
"...ij,...jk->...ik",
_dreamx_invert_se3(viewmats),
_dreamx_lift_k(_dreamx_invert_k(intrinsics_norm)),
).to(dtype=viewmats.dtype)
q = _dreamx_apply_tiled_projmat(q, proj_t)
k = _dreamx_apply_tiled_projmat(k, proj_inv)
v = _dreamx_apply_tiled_projmat(v, proj_inv)
return q, k, v, proj
class DreamXPropeSelfAttention(nn.Module):
"""DreamX-World parallel PRoPE camera self-attention branch."""
def __init__(self,
dim: int,
attn_dim: int,
num_heads: int,
qk_norm: str | bool = True,
eps: float = 1e-6,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
assert attn_dim % num_heads == 0
self.attn_dim = attn_dim
self.num_heads = num_heads
self.head_dim = attn_dim // num_heads
self.qk_norm = qk_norm
self.q_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.q_proj")
self.k_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.k_proj")
self.v_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.v_proj")
self.out_proj = ReplicatedLinear(attn_dim,
dim,
quant_config=quant_config,
prefix=f"{prefix}.out_proj")
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(self.head_dim, eps=eps)
self.norm_k = RMSNorm(self.head_dim, eps=eps)
elif qk_norm in (True, "rms_norm_across_heads"):
self.norm_q = RMSNorm(attn_dim, eps=eps)
self.norm_k = RMSNorm(attn_dim, eps=eps)
elif qk_norm is False:
self.norm_q = nn.Identity()
self.norm_k = nn.Identity()
else:
raise ValueError(f"Unsupported qk_norm for DreamX PRoPE: {qk_norm}")
nn.init.zeros_(self.out_proj.weight)
if self.out_proj.bias is not None:
nn.init.zeros_(self.out_proj.bias)
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self, hidden_states: torch.Tensor,
y_camera: dict[str, torch.Tensor]) -> torch.Tensor:
if get_sp_world_size() > 1:
# The transformer shards the sequence before the block loop and
# this branch uses LocalAttention (no all-to-all): under
# sequence parallelism each rank would attend only within its
# own shard — silently wrong output. Fail loudly until this
# path is ported to DistributedAttention and validated.
raise NotImplementedError(
"DreamXPropeSelfAttention does not support sequence "
"parallelism yet (LocalAttention on a sharded sequence "
"corrupts output). Run with sp_size=1.")
batch_size, seq_len, _ = hidden_states.shape
query, _ = self.q_proj(hidden_states)
key, _ = self.k_proj(hidden_states)
value, _ = self.v_proj(hidden_states)
if self.qk_norm == "rms_norm":
query = query.view(batch_size, seq_len, self.num_heads,
self.head_dim)
key = key.view(batch_size, seq_len, self.num_heads, self.head_dim)
query = self.norm_q(query)
key = self.norm_k(key)
else:
query = self.norm_q(query).view(batch_size, seq_len,
self.num_heads, self.head_dim)
key = self.norm_k(key).view(batch_size, seq_len, self.num_heads,
self.head_dim)
value = value.view(batch_size, seq_len, self.num_heads, self.head_dim)
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
query, key, value, output_projection = _dreamx_prope_qkv(
query,
key,
value,
viewmats=y_camera["viewmats"],
intrinsics=y_camera["K"],
)
out = self.attn(query.transpose(1, 2), key.transpose(1, 2),
value.transpose(1, 2))
out = _dreamx_apply_tiled_projmat(out.transpose(1, 2),
output_projection).transpose(1, 2)
out = out.flatten(2)
out, _ = self.out_proj(out)
return out
class DreamXWorldTransformerBlock(WanTransformerBlock):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
add_control_adapter: bool = True,
cam_method: str | None = "prope",
attn_compress: int = 1,
cam_self_attn_layers: tuple[int, ...] | None = None,
layer_idx: int | None = None):
super().__init__(dim, ffn_dim, num_heads, qk_norm, cross_attn_norm,
eps, added_kv_proj_dim,
supported_attention_backends, quant_config, prefix)
self.cam_self_attn = None
add_cam_attn = add_control_adapter and cam_method == "prope"
if add_cam_attn and cam_self_attn_layers is not None:
add_cam_attn = layer_idx in cam_self_attn_layers
if add_cam_attn:
if num_heads % attn_compress != 0 or dim % attn_compress != 0:
raise ValueError("DreamX attn_compress must divide dim and num_heads")
self.cam_self_attn = DreamXPropeSelfAttention(
dim,
dim // attn_compress,
num_heads // attn_compress,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.cam_self_attn")
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
original_seq_len: int,
y_camera: dict[str, torch.Tensor] | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
orig_dtype = hidden_states.dtype
if temb.dim() == 4:
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()).chunk(
6, dim=2)
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
attn_output, _ = self.attn1(
query,
key,
value,
original_seq_len,
freqs_cis=freqs_cis,
)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
if self.cam_self_attn is not None and y_camera is not None:
attn_output = attn_output + self.cam_self_attn(
norm_hidden_states, y_camera)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class DreamXWorldTransformer3DModel(WanTransformer3DModel):
_fsdp_shard_conditions = DreamXWorldConfig()._fsdp_shard_conditions
_compile_conditions = DreamXWorldConfig()._compile_conditions
_supported_attention_backends = DreamXWorldConfig(
)._supported_attention_backends
param_names_mapping = DreamXWorldConfig().param_names_mapping
reverse_param_names_mapping = DreamXWorldConfig().reverse_param_names_mapping
lora_param_names_mapping = DreamXWorldConfig().lora_param_names_mapping
def __init__(self, config: DreamXWorldConfig, hf_config: dict[str,
Any]) -> None:
BaseDiT.__init__(self, config=config, hf_config=hf_config)
self.quant_config = config.quant_config
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
self.blocks = nn.ModuleList([
DreamXWorldTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
quant_config=config.quant_config,
prefix=f"{config.prefix}.blocks.{i}",
add_control_adapter=config.add_control_adapter,
cam_method=config.cam_method,
attn_compress=config.attn_compress,
cam_self_attn_layers=config.cam_self_attn_layers,
layer_idx=i)
for i in range(config.num_layers)
])
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
y_camera: dict[str, torch.Tensor] | None = None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if encoder_hidden_states is not None and not isinstance(
encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, _, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
freqs_sin.to(hidden_states.device).float())
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
hidden_states, original_seq_len = sequence_model_parallel_shard(
hidden_states, dim=1)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten()
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep,
encoder_hidden_states,
encoder_hidden_states_image,
timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
assert encoder_hidden_states.dtype == orig_dtype
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, original_seq_len, y_camera)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, original_seq_len,
y_camera=y_camera)
if temb.dim() == 3:
shift, scale = (self.scale_shift_table.unsqueeze(0) +
temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
shift, scale = (self.scale_shift_table +
temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = sequence_model_parallel_all_gather_with_unpad(
hidden_states, original_seq_len, dim=1)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
EntryClass = DreamXWorldTransformer3DModel
+920
View File
@@ -0,0 +1,920 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World autoregressive causal DiT.
Adapted from DreamX-World's Apache-2.0
``wan/modules/causal_camera_model_2_2_prope_infinity.py``. The implementation is
kept native to FastVideo: no production import from DreamX, Diffusers, or
Transformers is required.
"""
from typing import Any
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.dreamx_world import (_dreamx_apply_tiled_projmat,
_dreamx_prope_qkv)
def attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
# Deliberately raw SDPA rather than fastvideo.attention.LocalAttention:
# (1) LocalAttention dispatches through the attention-backend registry, so
# FLASH_ATTN could be selected and its kernel is not bit-identical to
# torch SDPA — the AR KV-cache rollout must stay numerically frozen;
# (2) LocalAttention requires an active ForwardContext, which direct
# transformer invocations (parity tests) do not set;
# (3) the sibling causal model keeps raw SDPA in the same KV-cache window
# path (matrixgame2/causal_model.py).
# Sequence-parallel gap: this model never shards the sequence; run with
# sp_size=1 (see fastvideo/layers/AGENTS.md on documenting raw SDPA).
q_bhld = q.transpose(1, 2)
k_bhld = k.transpose(1, 2)
v_bhld = v.transpose(1, 2)
out = F.scaled_dot_product_attention(q_bhld, k_bhld, v_bhld, dropout_p=0.0)
return out.transpose(1, 2)
def prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
viewmats: torch.Tensor, Ks: torch.Tensor):
q, k, v, output_projection = _dreamx_prope_qkv(q, k, v, viewmats, Ks)
def apply_fn_o(x: torch.Tensor) -> torch.Tensor:
return _dreamx_apply_tiled_projmat(x, output_projection)
return q, k, v, apply_fn_o
def sinusoidal_embedding_1d(dim, position):
assert dim % 2 == 0
half = dim // 2
position = position.type(torch.float64)
sinusoid = torch.outer(
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
def rope_params(max_seq_len, dim, theta=10000):
assert dim % 2 == 0
freqs = torch.outer(
torch.arange(max_seq_len),
1.0 / torch.pow(theta,
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
return torch.polar(torch.ones_like(freqs), freqs)
class WanRMSNorm(nn.Module):
"""Kept private instead of fastvideo.layers.layernorm.RMSNorm.
The official DreamX-World ``model_2_2.py`` computes the RMS statistics in
the *input* dtype — the upstream code has the fp32 upcast explicitly
commented out (``# return self._norm(x.float())...``). FastVideo's RMSNorm
always normalizes in fp32, which is not bit-identical under bf16, so the
verbatim implementation stays.
"""
def __init__(self, dim, eps=1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
return self._norm(x).type_as(x) * self.weight
def _norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
class WanLayerNorm(nn.LayerNorm):
"""Kept private instead of fastvideo.layers.layernorm.FP32LayerNorm.
The official DreamX-World ``model_2_2.py`` normalizes in the *input* dtype
(no ``x.float()`` upcast, unlike Wan2.1). FP32LayerNorm casts input and
affine params to fp32, which is not bit-identical under bf16, so the
verbatim implementation stays.
"""
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
def forward(self, x):
return super().forward(x).type_as(x)
class WanCrossAttention(nn.Module):
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6):
super().__init__()
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = ReplicatedLinear(dim, dim)
self.k = ReplicatedLinear(dim, dim)
self.v = ReplicatedLinear(dim, dim)
self.o = ReplicatedLinear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def _kv(self, context, b, n, d):
k, _ = self.k(context)
k = self.norm_k(k).view(b, -1, n, d)
v, _ = self.v(context)
v = v.view(b, -1, n, d)
return k, v
def forward(self, x, context, context_lens, crossattn_cache=None):
b, n, d = x.size(0), self.num_heads, self.head_dim
q, _ = self.q(x)
q = self.norm_q(q).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k, v = self._kv(context, b, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k, v = self._kv(context, b, n, d)
x = attention(q, k, v)
x = x.flatten(2)
out, _ = self.o(x)
return out
def block_relativistic_rope(x, grid_sizes, freqs, start_frame=0, relative_frame_indices=None):
"""
Apply Block-Relativistic RoPE to input tensor.
Adapted from Infinity-RoPE (https://arxiv.org/abs/2511.20649).
Args:
x: Input tensor [B, L, num_heads, head_dim]
grid_sizes: Tensor [B, 3] containing (F, H, W)
freqs: RoPE frequencies
start_frame: Starting frame index for sequential RoPE
relative_frame_indices: Optional tensor [F] specifying explicit frame indices
for Block-Relativistic RoPE. Overrides start_frame if provided.
"""
n, c = x.size(2), x.size(3) // 2
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
seq_len, n, -1, 2))
if relative_frame_indices is not None:
frame_indices = relative_frame_indices.long()
freqs_temporal = freqs[0][frame_indices].view(f, 1, 1, -1).expand(f, h, w, -1)
else:
freqs_temporal = freqs[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1)
freqs_i = torch.cat([
freqs_temporal,
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(seq_len, 1, -1)
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
output.append(x_i)
return torch.stack(output).type_as(x)
class CausalWanSelfAttention(nn.Module):
"""Self-attention with KV cache and Block-Relativistic RoPE for causal inference."""
def __init__(self, dim, num_heads, local_attn_size=6, sink_size=1,
qk_norm=True, eps=1e-6):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
self.q = ReplicatedLinear(dim, dim)
self.k = ReplicatedLinear(dim, dim)
self.v = ReplicatedLinear(dim, dim)
self.o = ReplicatedLinear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, seq_lens, grid_sizes, freqs, kv_cache,
current_start=0, cache_start=None, sink_recache_after_switch=False):
"""
Args:
x: Shape [B, L, C]
seq_lens: Shape [B]
grid_sizes: Shape [B, 3] containing (F, H, W)
freqs: RoPE frequencies [1024, head_dim / 2]
kv_cache: Dict with 'k', 'v', 'global_end_index', 'local_end_index'
current_start: Current position in the global token sequence
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
if cache_start is None:
cache_start = current_start
q, _ = self.q(x)
q = self.norm_q(q).view(b, s, n, d)
k, _ = self.k(x)
k = self.norm_k(k).view(b, s, n, d)
v, _ = self.v(x)
v = v.view(b, s, n, d)
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
num_new_frames = grid_sizes[0][0].item()
current_end = current_start + q.shape[1]
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = q.shape[1]
cache_update_info = None
is_recompute = current_end <= kv_cache["global_end_index"].item() and current_start > 0
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# === ROLLING MODE: cache full, evict oldest non-sink tokens ===
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
temp_k = kv_cache["k"].detach().clone()
temp_v = kv_cache["v"].detach().clone()
temp_k[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
temp_k[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
temp_v[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
temp_v[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
# Block-Relativistic RoPE: query uses window-relative indices
query_relative_indices = torch.arange(
self.local_attn_size - num_new_frames, self.local_attn_size, device=q.device)
roped_query = block_relativistic_rope(
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
# Block-Relativistic RoPE: cached K uses position-in-window indices
num_cache_frames = local_end_index // frame_seqlen
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
cache_grid_sizes = grid_sizes.clone()
cache_grid_sizes[0, 0] = num_cache_frames
roped_temp_k = block_relativistic_rope(
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
cache_update_info = {
"action": "roll_and_insert",
"sink_tokens": sink_tokens,
"num_rolled_tokens": num_rolled_tokens,
"num_evicted_tokens": num_evicted_tokens,
"local_start_index": local_start_index,
"local_end_index": local_end_index,
"write_start_index": write_start_index,
"write_end_index": local_end_index,
"new_k": k[:, roped_offset:roped_offset + write_len],
"new_v": v[:, roped_offset:roped_offset + write_len],
"current_end": current_end,
"is_recompute": is_recompute
}
else:
# === DIRECT INSERT MODE: cache not yet full ===
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
temp_k = kv_cache["k"].detach().clone()
temp_v = kv_cache["v"].detach().clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
if sink_recache_after_switch:
write_start_index = local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
# RoPE with relative indices (growing sequentially before cache fills)
current_frame_in_window = local_start_index // frame_seqlen
query_relative_indices = torch.arange(
current_frame_in_window, current_frame_in_window + num_new_frames, device=q.device)
roped_query = block_relativistic_rope(
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
num_cache_frames = local_end_index // frame_seqlen
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
cache_grid_sizes = grid_sizes.clone()
cache_grid_sizes[0, 0] = num_cache_frames
roped_temp_k = block_relativistic_rope(
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
cache_update_info = {
"action": "direct_insert",
"local_start_index": local_start_index,
"local_end_index": local_end_index,
"write_start_index": write_start_index,
"write_end_index": local_end_index,
"new_k": k[:, roped_offset:roped_offset + write_len],
"new_v": v[:, roped_offset:roped_offset + write_len],
"current_end": current_end,
"is_recompute": is_recompute
}
# Attention: sink tokens + local window
if sink_tokens > 0:
local_budget = self.max_attention_size - sink_tokens
k_sink = roped_temp_k[:, :sink_tokens]
v_sink = temp_v[:, :sink_tokens]
if local_budget > 0:
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
k_local = roped_temp_k[:, local_start_for_window:local_end_index]
v_local = temp_v[:, local_start_for_window:local_end_index]
k_cat = torch.cat([k_sink, k_local], dim=1)
v_cat = torch.cat([v_sink, v_local], dim=1)
else:
k_cat = k_sink
v_cat = v_sink
x = attention(roped_query, k_cat, v_cat)
else:
window_start = max(0, local_end_index - self.max_attention_size)
x = attention(
roped_query,
roped_temp_k[:, window_start:local_end_index],
temp_v[:, window_start:local_end_index])
x = x.flatten(2)
x, _ = self.o(x)
return x, (current_end, local_end_index, cache_update_info)
class CausalPropeSelfAttention(nn.Module):
"""PRoPE self-attention with optional KV cache for camera-controlled inference."""
def __init__(self, dim, attn_dim, num_heads, window_size=(-1, -1),
local_attn_size=-1, sink_size=0, qk_norm=True, eps=1e-6):
assert dim % num_heads == 0
assert attn_dim % num_heads == 0
super().__init__()
self.dim = dim
self.attn_dim = attn_dim
self.num_heads = num_heads
self.head_dim = attn_dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.window_size = window_size
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
self.q_proj = ReplicatedLinear(dim, attn_dim)
self.k_proj = ReplicatedLinear(dim, attn_dim)
self.v_proj = ReplicatedLinear(dim, attn_dim)
self.out_proj = ReplicatedLinear(attn_dim, dim)
self.norm_q = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
nn.init.zeros_(self.out_proj.weight)
nn.init.zeros_(self.out_proj.bias)
def forward(self, x, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
kv_cache=None, current_start=0, cache_start=None,
sink_recache_after_switch=False, cache_update_policy="commit_detached"):
"""
Args:
x: Shape [B, L, C]
cam_viewmats: Camera view matrices
cam_K: Camera intrinsics
kv_cache: Optional KV cache dict. When None, runs full attention over current chunk.
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
if cache_start is None:
cache_start = current_start
q, _ = self.q_proj(x)
q = self.norm_q(q).view(b, s, n, d)
k, _ = self.k_proj(x)
k = self.norm_k(k).view(b, s, n, d)
v, _ = self.v_proj(x)
v = v.view(b, s, n, d)
# Apply PRoPE (Positional Rotary Position Embedding from camera parameters)
q_t, k_t, v_t, apply_fn_o = prope_qkv(
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2),
viewmats=cam_viewmats, Ks=cam_K)
proped_q = q_t.transpose(1, 2)
proped_k = k_t.transpose(1, 2)
proped_v = v_t.transpose(1, 2)
if kv_cache is None:
# No cache: full attention over current chunk
x_out = attention(proped_q, proped_k, proped_v)
else:
# KV cache mode with rolling cache support
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
num_new_tokens = s
current_end = current_start + num_new_tokens
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
is_recompute = (current_end <= kv_cache["global_end_index"].item()) and (current_start > 0)
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# === ROLLING MODE ===
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
if cache_update_policy != "none":
with torch.no_grad():
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
with torch.no_grad():
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
else:
# === DIRECT INSERT MODE ===
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
if cache_update_policy != "none":
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
if sink_recache_after_switch:
write_start_index = local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
with torch.no_grad():
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
# Attention: sink tokens + local window
if sink_tokens > 0:
local_budget = self.max_attention_size - sink_tokens
k_sink = kv_cache["k"][:, :sink_tokens].detach()
v_sink = kv_cache["v"][:, :sink_tokens].detach()
if local_budget > 0:
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
k_local = kv_cache["k"][:, local_start_for_window:local_end_index].detach()
v_local = kv_cache["v"][:, local_start_for_window:local_end_index].detach()
k_cat = torch.cat([k_sink, k_local], dim=1)
v_cat = torch.cat([v_sink, v_local], dim=1)
else:
k_cat = k_sink
v_cat = v_sink
x_out = attention(proped_q, k_cat, v_cat)
else:
window_start = max(0, local_end_index - self.max_attention_size)
x_out = attention(
proped_q,
kv_cache["k"][:, window_start:local_end_index].detach(),
kv_cache["v"][:, window_start:local_end_index].detach())
if not is_recompute and cache_update_policy != "none":
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
# Apply inverse PRoPE
x = apply_fn_o(x_out.transpose(1, 2)).transpose(1, 2)
x = x.flatten(2)
x, _ = self.out_proj(x)
return x
class CausalWanAttentionBlock(nn.Module):
def __init__(self, dim, ffn_dim, num_heads, local_attn_size=-1, sink_size=0,
qk_norm=True, cross_attn_norm=False, eps=1e-6, **kwargs):
super().__init__()
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.local_attn_size = local_attn_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.add_control_adapter = kwargs.get('add_control_adapter', False)
self.cam_method = kwargs.get('cam_method')
self.attn_compress = kwargs.get('attn_compress', 1)
self.layer_idx = kwargs.get('layer_idx')
cam_self_attn_layers = kwargs.get('cam_self_attn_layers')
# layers
self.norm1 = WanLayerNorm(dim, eps)
self.self_attn = CausalWanSelfAttention(
dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
self.norm3 = WanLayerNorm(
dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.cross_attn = WanCrossAttention(dim, num_heads, (-1, -1), qk_norm, eps)
self.norm2 = WanLayerNorm(dim, eps)
# nn.Linear (not ReplicatedLinear) on purpose: the official checkpoint
# stores these as positional Sequential keys (ffn.0 / ffn.2) that the
# copy-only converter and the strict-load tests require verbatim, and
# ReplicatedLinear's (out, bias) tuple return cannot compose inside
# nn.Sequential without renaming the state-dict surface.
self.ffn = nn.Sequential(
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
nn.Linear(ffn_dim, dim))
# PRoPE self-attention branch for camera control
add_cam_attn = self.add_control_adapter and self.cam_method == 'prope'
if add_cam_attn and cam_self_attn_layers is not None:
add_cam_attn = self.layer_idx in cam_self_attn_layers
if add_cam_attn:
self.cam_self_attn = CausalPropeSelfAttention(
dim, dim // self.attn_compress, num_heads,
local_attn_size=local_attn_size, sink_size=sink_size,
qk_norm=qk_norm, eps=eps)
# modulation
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(self, x, e, seq_lens, grid_sizes, freqs, context, context_lens,
kv_cache, crossattn_cache=None, current_start=0, cache_start=None,
cam_viewmats=None, cam_K=None, sink_recache_after_switch=False,
cache_update_policy="commit_detached"):
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
# self-attention
attn_input = (self.norm1(x).unflatten(
dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2)
y, cache_update_info = self.self_attn(
attn_input, seq_lens, grid_sizes, freqs, kv_cache,
current_start, cache_start, sink_recache_after_switch)
# PRoPE camera attention (parallel branch)
if hasattr(self, 'cam_self_attn') and cam_viewmats is not None and cam_K is not None:
prope_kv_cache = None
if kv_cache is not None and "prope_k" in kv_cache:
prope_kv_cache = {
"k": kv_cache["prope_k"],
"v": kv_cache["prope_v"],
"global_end_index": kv_cache["prope_global_end_index"],
"local_end_index": kv_cache["prope_local_end_index"],
}
y = y + self.cam_self_attn(
attn_input, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
kv_cache=prope_kv_cache, current_start=current_start,
cache_start=cache_start, cache_update_policy=cache_update_policy)
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2)
# cross-attention & FFN
x = x + self.cross_attn(self.norm3(x), context, context_lens,
crossattn_cache=crossattn_cache)
y = self.ffn(
(self.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
* (1 + e[4]) + e[3]).flatten(1, 2))
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]).flatten(1, 2)
return x, cache_update_info
class CausalHead(nn.Module):
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
super().__init__()
self.dim = dim
self.out_dim = out_dim
self.patch_size = patch_size
self.eps = eps
out_dim = math.prod(patch_size) * out_dim
self.norm = WanLayerNorm(dim, eps)
self.head = ReplicatedLinear(dim, out_dim)
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
def forward(self, x, e):
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
e = (self.modulation.unsqueeze(1) + e).chunk(2, dim=2)
x, _ = self.head(
self.norm(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
* (1 + e[1]) + e[0])
return x
class DreamXWorldARTransformer3DModel(BaseDiT):
"""DreamX-World-5B autoregressive causal transformer."""
_fsdp_shard_conditions = DreamXWorldARConfig()._fsdp_shard_conditions
_compile_conditions = DreamXWorldARConfig()._compile_conditions
_supported_attention_backends = DreamXWorldARConfig()._supported_attention_backends
param_names_mapping = DreamXWorldARConfig().param_names_mapping
reverse_param_names_mapping = DreamXWorldARConfig().reverse_param_names_mapping
lora_param_names_mapping = DreamXWorldARConfig().lora_param_names_mapping
_no_split_modules = ["CausalWanAttentionBlock"]
def __init__(self, config: DreamXWorldARConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
model_type = config.model_type
patch_size = config.patch_size
text_len = config.text_len
in_dim = config.in_channels
dim = config.hidden_size
ffn_dim = config.ffn_dim
freq_dim = config.freq_dim
text_dim = config.text_dim
out_dim = config.out_channels
num_heads = config.num_attention_heads
num_layers = config.num_layers
local_attn_size = config.local_attn_size
sink_size = config.sink_size
qk_norm = bool(config.qk_norm)
cross_attn_norm = config.cross_attn_norm
eps = config.eps
add_control_adapter = config.add_control_adapter
cam_method = config.cam_method
attn_compress = config.attn_compress
cam_self_attn_layers = config.cam_self_attn_layers
assert model_type in ['t2v', 'i2v', 'ti2v']
self.model_type = model_type
self.patch_size = patch_size
self.text_len = text_len
self.in_dim = in_dim
self.dim = dim
self.ffn_dim = ffn_dim
self.freq_dim = freq_dim
self.text_dim = text_dim
self.out_dim = out_dim
self.num_heads = num_heads
self.num_layers = num_layers
self.local_attn_size = local_attn_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
# embeddings — nn.Linear inside nn.Sequential on purpose: the official
# checkpoint keys are positional (text_embedding.0/.2, time_embedding.0/.2,
# time_projection.1) and must load verbatim (see ffn comment above).
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.text_embedding = nn.Sequential(
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
nn.Linear(dim, dim))
self.time_embedding = nn.Sequential(
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(
nn.SiLU(), nn.Linear(dim, dim * 6))
# transformer blocks
self.blocks = nn.ModuleList([
CausalWanAttentionBlock(
dim, ffn_dim, num_heads, local_attn_size, sink_size,
qk_norm, cross_attn_norm, eps,
add_control_adapter=add_control_adapter,
cam_method=cam_method,
attn_compress=attn_compress,
layer_idx=layer_idx,
cam_self_attn_layers=cam_self_attn_layers)
for layer_idx in range(num_layers)
])
for layer_idx, block in enumerate(self.blocks):
block.self_attn.layer_idx = layer_idx
block.self_attn.num_layers = self.num_layers
# head
self.head = CausalHead(dim, out_dim, patch_size, eps)
# RoPE frequencies
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
d = dim // num_heads
self.freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
], dim=1)
self.num_attention_heads = num_heads
self.attention_head_dim = dim // num_heads
self.hidden_size = dim
self.in_channels = in_dim
self.out_channels = out_dim
self.num_channels_latents = out_dim
self.init_weights()
self.num_frame_per_block = config.arch_config.num_frames_per_block
self.__post_init__()
def forward(self, x=None, t=None, context=None, seq_len=None, y=None, y_camera=None,
kv_cache=None, crossattn_cache=None, current_start=0,
cache_start=0, cache_update_policy="commit_detached",
hidden_states=None, encoder_hidden_states=None, timestep=None, **kwargs):
"""
Causal inference with KV caching.
See Algorithm 2 of CausVid (https://arxiv.org/abs/2412.07772).
Args:
x: List of input video tensors [C_in, F, H, W]
t: Timestep tensor [B, L]
context: List of text embeddings [L, C]
seq_len: Maximum sequence length for positional encoding
y: Optional conditional video inputs (I2V mode)
y_camera: Camera parameters dict {'viewmats': ..., 'K': ...}
kv_cache: List of KV cache dicts per transformer block
crossattn_cache: List of cross-attention cache dicts
current_start: Current position in global token sequence
cache_start: Cache start position
cache_update_policy: Cache update strategy ('commit_detached' or 'none')
Returns:
Stacked output tensors [B, C_out, F, H/8, W/8]
"""
if x is None and hidden_states is not None:
x = [sample for sample in hidden_states]
if t is None and timestep is not None:
t = timestep
if context is None and encoder_hidden_states is not None:
if isinstance(encoder_hidden_states, torch.Tensor):
context = [sample for sample in encoder_hidden_states]
else:
context = encoder_hidden_states
if seq_len is None:
if torch.is_tensor(t):
seq_len = int(t.shape[1]) if t.dim() > 1 else int(t.numel())
elif x is not None:
sample = x[0]
seq_len = (sample.shape[1] // self.patch_size[0]) * (sample.shape[2] // self.patch_size[1]) * (sample.shape[3] // self.patch_size[2])
if x is None or t is None or context is None or seq_len is None:
raise ValueError("DreamXWorldARTransformer3DModel requires x/t/context/seq_len or FastVideo aliases")
device = self.patch_embedding.weight.device
if self.freqs.is_meta or self.freqs.device != device:
d = self.dim // self.num_heads
self.freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
], dim=1).to(device)
if y is not None:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
# patch embedding
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
assert seq_lens.max() <= seq_len
x = torch.cat(x)
# time embedding
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x))
e0 = self.time_projection(e).unflatten(
1, (6, self.dim)).unflatten(dim=0, sizes=t.shape)
# text embedding
context_lens = None
context = self.text_embedding(
torch.stack([
torch.cat(
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]))
# camera parameters
if y_camera is not None and isinstance(y_camera, dict):
cam_viewmats = y_camera['viewmats']
cam_K = y_camera['K']
else:
cam_viewmats = None
cam_K = None
block_kwargs = dict(
e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=self.freqs,
context=context, context_lens=context_lens,
cam_viewmats=cam_viewmats, cam_K=cam_K,
cache_update_policy=cache_update_policy,
)
cache_update_infos = []
for block_index, block in enumerate(self.blocks):
block_kwargs.update({
"kv_cache": kv_cache[block_index] if kv_cache is not None else None,
"crossattn_cache": crossattn_cache[block_index] if crossattn_cache is not None else None,
"current_start": current_start,
"cache_start": cache_start,
})
x, block_cache_update_info = block(x, **block_kwargs)
if kv_cache is not None:
cache_update_infos.append((block_index, block_cache_update_info))
# Apply deferred cache updates
if kv_cache is not None and cache_update_infos and cache_update_policy != "none":
self._apply_cache_updates(kv_cache, cache_update_infos)
# head & unpatchify
x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2))
x = self.unpatchify(x, grid_sizes)
return torch.stack(x)
def _apply_cache_updates(self, kv_cache, cache_update_infos):
"""Apply deferred cache updates collected from all transformer blocks.
For Block-Relativistic RoPE, this stores un-roped K values in the cache.
RoPE is applied dynamically during attention based on each token's current
relative position in the sliding window.
"""
with torch.no_grad():
for block_index, (current_end, local_end_index, update_info) in cache_update_infos:
if update_info is not None:
cache = kv_cache[block_index]
if update_info["action"] == "roll_and_insert":
sink_tokens = update_info["sink_tokens"]
num_rolled_tokens = update_info["num_rolled_tokens"]
num_evicted_tokens = update_info["num_evicted_tokens"]
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
new_k = update_info["new_k"].detach()
new_v = update_info["new_v"].detach()
cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
cache["k"][:, write_start_index:write_end_index] = new_k
cache["v"][:, write_start_index:write_end_index] = new_v
elif update_info["action"] == "direct_insert":
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
new_k = update_info["new_k"].detach()
new_v = update_info["new_v"].detach()
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
cache["k"][:, write_start_index:write_end_index] = new_k
cache["v"][:, write_start_index:write_end_index] = new_v
is_recompute = False if update_info is None else update_info.get("is_recompute", False)
if not is_recompute:
kv_cache[block_index]["global_end_index"].fill_(current_end)
kv_cache[block_index]["local_end_index"].fill_(local_end_index)
def unpatchify(self, x, grid_sizes):
"""Reconstruct video tensors from patch embeddings."""
c = self.out_dim
out = []
for u, v in zip(x, grid_sizes.tolist(), strict=True):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = torch.einsum('fhwpqrc->cfphqwr', u)
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
out.append(u)
return out
def init_weights(self):
"""Initialize model parameters using Xavier initialization."""
for m in self.modules():
if isinstance(m, (nn.Linear, ReplicatedLinear)):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
for m in self.text_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
for m in self.time_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
nn.init.zeros_(self.head.head.weight)
EntryClass = DreamXWorldARTransformer3DModel
@@ -537,9 +537,9 @@ class CausalMatrixGame2TransformerBlock(nn.Module):
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
+7 -3
View File
@@ -11,7 +11,7 @@ import tempfile
from abc import ABC, abstractmethod
from collections.abc import Callable, Set
from dataclasses import dataclass, field
from functools import lru_cache
from functools import cache, lru_cache
from typing import NoReturn, TypeVar, cast
import cloudpickle
@@ -32,6 +32,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HYWorldTransformer3DModel":
("dits", "hyworld", "HYWorldTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
@@ -48,6 +50,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
@@ -141,7 +145,7 @@ _LEGACY_FAST_VIDEO_MODELS = {
MODELS_PATH = os.path.dirname(__file__)
@lru_cache(maxsize=None)
@cache
def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
discovered_models: dict[str, tuple[str, str, str]] = {}
for root, dirs, files in os.walk(MODELS_PATH):
@@ -156,7 +160,7 @@ def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
filepath = os.path.join(root, filename)
try:
with open(filepath, "r", encoding="utf-8") as f:
with open(filepath, encoding="utf-8") as f:
source = f.read()
tree = ast.parse(source, filename=filename)
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
"""Performance benchmark and dashboard utilities."""
@@ -57,9 +57,7 @@ def is_baseline_eligible_record(record: dict[str, Any]) -> bool:
"""
if record.get("baseline_eligible") is True:
return True
if "baseline_eligible" not in record and "run_source" not in record:
return True
return False
return "baseline_eligible" not in record and "run_source" not in record
def resolve_hf_token() -> str | None:
@@ -312,8 +310,12 @@ def load_records_for_model(
# ---------------------------------------------------------------------------
_NUMERIC_COLS = (
"latency", "throughput", "memory",
"text_encoder_time_s", "dit_time_s", "vae_decode_time_s",
"latency",
"throughput",
"memory",
"text_encoder_time_s",
"dit_time_s",
"vae_decode_time_s",
)
+116
View File
@@ -0,0 +1,116 @@
# SPDX-License-Identifier: Apache-2.0
"""Metric policy for rolling performance baseline comparisons."""
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any
@dataclass(frozen=True)
class MetricPolicy:
key: str
label: str
precision: int
lower_is_better: bool
threshold_percent: float
threshold_absolute: float
gated: bool = True
@dataclass(frozen=True)
class MetricDelta:
absolute: float
percent: float
threshold_exceeded: bool
regressed: bool
DEFAULT_METRIC_POLICIES: tuple[MetricPolicy, ...] = (
MetricPolicy("latency", "Latency", 3, True, 0.08, 0.5),
MetricPolicy("throughput", "Throughput", 3, False, 0.08, 0.05),
MetricPolicy("memory", "Memory", 1, True, 0.05, 256.0),
MetricPolicy("text_encoder_time_s", "Text Enc", 3, True, 0.05, 0.25),
MetricPolicy("dit_time_s", "DiT", 3, True, 0.05, 0.25),
MetricPolicy("vae_decode_time_s", "VAE Decode", 3, True, 0.05, 0.25),
)
def _optional_float(value: Any) -> float | None:
if value is None or isinstance(value, bool):
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def _optional_bool(value: Any) -> bool | None:
if isinstance(value, bool):
return value
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in {"1", "true", "yes", "on"}:
return True
if normalized in {"0", "false", "no", "off"}:
return False
return None
def resolve_metric_policies(threshold_overrides: Mapping[str, Any] | None, ) -> tuple[MetricPolicy, ...]:
"""Return default metric policies with optional per-metric overrides."""
if not isinstance(threshold_overrides, Mapping):
threshold_overrides = {}
policies: list[MetricPolicy] = []
for base_policy in DEFAULT_METRIC_POLICIES:
raw_override = threshold_overrides.get(base_policy.key, {})
if not isinstance(raw_override, Mapping):
raw_override = {}
threshold_percent = _optional_float(raw_override.get("threshold_percent"))
threshold_absolute = _optional_float(raw_override.get("threshold_absolute"))
gated = _optional_bool(raw_override.get("gated"))
policies.append(
MetricPolicy(
key=base_policy.key,
label=base_policy.label,
precision=base_policy.precision,
lower_is_better=base_policy.lower_is_better,
threshold_percent=(base_policy.threshold_percent if threshold_percent is None else threshold_percent),
threshold_absolute=(base_policy.threshold_absolute
if threshold_absolute is None else threshold_absolute),
gated=base_policy.gated if gated is None else gated,
))
return tuple(policies)
def serialize_metric_thresholds(policies: tuple[MetricPolicy, ...], ) -> dict[str, dict[str, float | bool]]:
return {
policy.key: {
"threshold_percent": policy.threshold_percent,
"threshold_absolute": policy.threshold_absolute,
"gated": policy.gated,
}
for policy in policies
}
def regression_delta(
policy: MetricPolicy,
current: float,
baseline: float,
) -> MetricDelta | None:
if baseline <= 0:
return None
absolute_delta = current - baseline if policy.lower_is_better else baseline - current
percent_delta = absolute_delta / baseline
threshold_exceeded = (percent_delta > policy.threshold_percent and absolute_delta > policy.threshold_absolute)
return MetricDelta(
absolute=absolute_delta,
percent=percent_delta,
threshold_exceeded=threshold_exceeded,
regressed=policy.gated and threshold_exceeded,
)
+1 -2
View File
@@ -13,7 +13,7 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles
from fastvideo.tests.performance import hf_store
from fastvideo.performance import hf_store
from .service import build_latest_summary, build_trends, filter_records
@@ -150,7 +150,6 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
rows = build_latest_summary(
filtered,
max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")),
run_source=run_source,
)
return {
+2 -20
View File
@@ -1,26 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
"""Metric definitions shared by the performance dashboard backend."""
from __future__ import annotations
from fastvideo.performance.metric_policy import DEFAULT_METRIC_POLICIES
from dataclasses import dataclass
@dataclass(frozen=True)
class MetricDefinition:
key: str
label: str
precision: int
lower_is_better: bool
METRICS: tuple[MetricDefinition, ...] = (
MetricDefinition("latency", "Latency", 3, True),
MetricDefinition("throughput", "Throughput", 3, False),
MetricDefinition("memory", "Memory", 1, True),
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
MetricDefinition("dit_time_s", "DiT", 3, True),
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
)
METRICS = DEFAULT_METRIC_POLICIES
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
+41 -45
View File
@@ -13,9 +13,8 @@ from collections import defaultdict
from datetime import datetime, timezone
from typing import Any
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
from .metrics import METRICS
from fastvideo.performance.hf_store import is_baseline_eligible_record, safe_float
from fastvideo.performance.metric_policy import regression_delta, resolve_metric_policies
Record = dict[str, Any]
@@ -95,19 +94,9 @@ def baseline_value(records: list[Record], metric_key: str) -> float | None:
return float(statistics.median(values))
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
if current is None or baseline is None or baseline <= 0:
return None
metric = next(metric for metric in METRICS if metric.key == metric_key)
if metric.lower_is_better:
return (current - baseline) / baseline * 100.0
return (baseline - current) / baseline * 100.0
def build_latest_summary(records: list[Record],
*,
baseline_window: int = 5,
max_regression: float = 0.05,
run_source: str | None = None) -> list[Record]:
rows: list[Record] = []
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
@@ -123,52 +112,58 @@ def build_latest_summary(records: list[Record],
if record is not latest and record.get("success", True) and is_baseline_eligible_record(record)
]
baseline_records = baseline_pool[-baseline_window:]
metric_policies = resolve_metric_policies(latest.get("regression_thresholds"))
metrics: dict[str, Record] = {}
regressions: list[float] = []
for metric in METRICS:
current = safe_float(latest.get(metric.key))
baseline = baseline_value(baseline_records, metric.key)
regression = regression_percent(metric.key, current, baseline)
metrics[metric.key] = {
failing_metrics: list[str] = []
threshold_exceeded_metrics: list[str] = []
for policy in metric_policies:
current = safe_float(latest.get(policy.key))
baseline = baseline_value(baseline_records, policy.key)
delta = None
if current is not None and baseline is not None:
delta = regression_delta(policy, current, baseline)
regression = None if delta is None else delta.percent * 100.0
metrics[policy.key] = {
"current": current,
"baseline": baseline,
"regression_pct": regression,
"label": metric.label,
"lower_is_better": metric.lower_is_better,
"precision": metric.precision,
"absolute_delta": None if delta is None else delta.absolute,
"threshold_percent": policy.threshold_percent * 100.0,
"threshold_absolute": policy.threshold_absolute,
"gated": policy.gated,
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
"regressed": False if delta is None else delta.regressed,
"label": policy.label,
"lower_is_better": policy.lower_is_better,
"precision": policy.precision,
}
if regression is not None:
regressions.append(regression)
if delta is not None and delta.threshold_exceeded:
threshold_exceeded_metrics.append(policy.key)
if delta is not None and delta.regressed:
failing_metrics.append(policy.key)
worst_regression = max(regressions) if regressions else None
success = bool(latest.get("success", True))
status = "pass" if success else "fail"
rows.append({
"model_id":
model_id,
"gpu_type":
gpu_type,
"timestamp":
latest.get("timestamp"),
"commit_sha":
latest.get("commit_sha"),
"model_id": model_id,
"gpu_type": gpu_type,
"timestamp": latest.get("timestamp"),
"commit_sha": latest.get("commit_sha"),
**record_metadata(latest),
"success":
success,
"baseline_n":
len(baseline_records),
"worst_regression_pct":
worst_regression,
"regression_threshold_pct":
max_regression * 100.0,
"computed_regression_status":
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
"status":
status,
"metrics":
metrics,
"success": success,
"baseline_n": len(baseline_records),
"worst_regression_pct": worst_regression,
"threshold_exceeded_metrics": threshold_exceeded_metrics,
"failing_metrics": failing_metrics,
"computed_regression_status": "fail" if failing_metrics else "pass",
"status": status,
"metrics": metrics,
})
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
@@ -179,14 +174,15 @@ def build_trends(records: list[Record]) -> list[Record]:
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
points = []
for record in group:
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
point = {
"timestamp": record.get("timestamp"),
"commit_sha": record.get("commit_sha"),
**record_metadata(record),
"success": bool(record.get("success", True)),
"metrics": {
metric.key: safe_float(record.get(metric.key))
for metric in METRICS
policy.key: safe_float(record.get(policy.key))
for policy in metric_policies
},
}
points.append(point)
@@ -0,0 +1,28 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BARPipelineConfig,
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_ar_dit_config,
make_dreamx_world_5b_cam_dit_config,
make_dreamx_world_5b_cam_text_encoder_config,
make_dreamx_world_5b_cam_vae_config,
)
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
__all__ = [
"DREAMX_Y_CAMERA_KEY",
"DreamXWorld5BARPipelineConfig",
"DreamXWorld5BCamPipelineConfig",
"DreamXWorldCameraConditioningStage",
"DreamXWorldARPipeline",
"DreamXWorldPipeline",
"make_dreamx_world_5b_ar_dit_config",
"make_dreamx_world_5b_cam_dit_config",
"make_dreamx_world_5b_cam_text_encoder_config",
"make_dreamx_world_5b_cam_vae_config",
]
@@ -0,0 +1,219 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World autoregressive causal denoising stage."""
from __future__ import annotations
from typing import Any
import torch
from tqdm.auto import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.basic.dreamx_world.stages import DREAMX_Y_CAMERA_KEY
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
class DreamXWorldARCausalDenoisingStage(DenoisingStage):
"""Official DreamX AR-forcing denoising loop with KV cache."""
_AR_NOISE_SEED_OFFSET = 1_000_003
def __init__(self, transformer, scheduler, pipeline=None, vae=None) -> None:
super().__init__(transformer=transformer, scheduler=scheduler, pipeline=pipeline, vae=vae)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frame_per_block = int(getattr(self.transformer, "num_frame_per_block", 3))
self.local_attn_size = int(getattr(self.transformer, "local_attn_size", 12))
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None, "latents must be prepared before DreamX AR denoising"
assert batch.prompt_embeds, "prompt embeds must be prepared before DreamX AR denoising"
latents = batch.latents
device = latents.device
target_dtype = torch.bfloat16
autocast_enabled = device.type == "cuda" and not fastvideo_args.disable_autocast
frame_seq_length = (latents.shape[-2] // self.transformer.patch_size[1]) * (latents.shape[-1] //
self.transformer.patch_size[2])
timesteps = torch.tensor(
tuple(getattr(fastvideo_args.pipeline_config, "dmd_denoising_steps", (1000, 750, 500, 250))),
dtype=torch.long,
).cpu()
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step", True):
self.scheduler.set_timesteps(1000)
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(device)
if latents.shape[2] % self.num_frame_per_block != 0:
raise ValueError("DreamX AR latent frames must be divisible by num_frame_per_block")
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY, batch.extra.get("y_camera"))
if isinstance(y_camera, dict):
y_camera = {
k: v.to(device=device, dtype=target_dtype) if torch.is_tensor(v) else v
for k, v in y_camera.items()
}
if batch.image_latent is not None and batch.image_latent.shape[1] == latents.shape[1]:
latents[:, :, :batch.image_latent.shape[2]] = batch.image_latent.to(device=device, dtype=latents.dtype)
kv_cache = self._initialize_kv_cache(latents.shape[0], target_dtype, device, frame_seq_length)
crossattn_cache = self._initialize_crossattn_cache(latents.shape[0], target_dtype, device)
prompt = batch.prompt_embeds[0]
if torch.is_tensor(prompt):
prompt = prompt.to(device=device, dtype=target_dtype)
context = [sample for sample in prompt]
else:
context = prompt
num_blocks = latents.shape[2] // self.num_frame_per_block
start = 0
first_frame_mask = torch.ones_like(latents)
first_frame_mask[:, :, 0] = 0
base_generator = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
noise_generator = self._make_noise_generator(base_generator, device)
with tqdm(total=num_blocks * len(timesteps), desc="DreamX AR denoising", leave=False) as progress:
for _ in range(num_blocks):
current_num_frames = self.num_frame_per_block
block_latents = latents[:, :, start:start + current_num_frames]
noisy_input = block_latents.clone()
mask_block = first_frame_mask[:, :, start:start + current_num_frames]
camera_block = self._slice_camera(y_camera, start, current_num_frames)
for idx, current_timestep in enumerate(timesteps):
timestep = torch.full(
(latents.shape[0], current_num_frames * frame_seq_length),
int(current_timestep.item()),
device=device,
dtype=torch.long,
)
if start == 0:
timestep[:, :frame_seq_length] = 0
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
denoised = self.transformer(
hidden_states=block_latents.to(target_dtype),
encoder_hidden_states=torch.stack(context).to(target_dtype),
timestep=timestep,
y_camera=camera_block,
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=start * frame_seq_length,
)
denoised = denoised.to(latents.dtype)
if idx < len(timesteps) - 1:
next_timestep = torch.full((latents.shape[0], current_num_frames),
int(timesteps[idx + 1].item()),
device=device,
dtype=torch.long)
noise_kwargs = {"device": device, "dtype": denoised.dtype}
if noise_generator is not None:
noise_kwargs["generator"] = noise_generator
noise = torch.randn(denoised.permute(0, 2, 1, 3, 4).shape, **noise_kwargs)
block_btchw = self.scheduler.add_noise(
denoised.permute(0, 2, 1, 3, 4).flatten(0, 1),
noise.flatten(0, 1),
next_timestep.flatten(),
).unflatten(0, (latents.shape[0], current_num_frames))
block_latents = block_btchw.permute(0, 2, 1, 3, 4)
block_latents = block_latents * mask_block + noisy_input * (1 - mask_block)
else:
block_latents = denoised * mask_block + noisy_input * (1 - mask_block)
progress.update()
latents[:, :, start:start + current_num_frames] = block_latents
self._update_context_cache(block_latents, context, camera_block, kv_cache, crossattn_cache, start,
frame_seq_length, target_dtype, autocast_enabled,
float(getattr(fastvideo_args.pipeline_config, "context_noise", 0.1)))
start += current_num_frames
batch.latents = latents
return batch
def _make_noise_generator(self, generator: torch.Generator | None, device: torch.device) -> torch.Generator | None:
if generator is None:
return None
if getattr(generator, "device", None) == device:
return generator
seed = int(generator.initial_seed()) + self._AR_NOISE_SEED_OFFSET
return torch.Generator(device=device).manual_seed(seed)
@staticmethod
def _context_noise_timestep(context_noise: float) -> int:
if 0.0 < context_noise <= 1.0:
return int(context_noise * 1000)
return int(context_noise)
def _slice_camera(self, y_camera: Any, start: int, num_frames: int):
if not isinstance(y_camera, dict):
return y_camera
return {
"viewmats": y_camera["viewmats"][:, start:start + num_frames],
"K": y_camera["K"][:, start:start + num_frames],
}
def _initialize_kv_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device,
frame_seq_length: int) -> list[dict[str, Any]]:
size = self.local_attn_size * frame_seq_length if self.local_attn_size != -1 else 18480
heads = self.transformer.num_attention_heads
head_dim = self.transformer.attention_head_dim
cam_self_attn = next(
(getattr(block, "cam_self_attn", None)
for block in self.transformer.blocks if getattr(block, "cam_self_attn", None) is not None),
None,
)
caches = []
for _ in range(self.num_transformer_blocks):
cache = {
"k": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
"v": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device),
}
if cam_self_attn is not None:
cam_heads = int(cam_self_attn.num_heads)
cam_head_dim = int(cam_self_attn.head_dim)
cache.update({
"prope_k":
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
"prope_v":
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
"prope_global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"prope_local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
caches.append(cache)
return caches
def _initialize_crossattn_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device):
heads = self.transformer.num_attention_heads
head_dim = self.transformer.attention_head_dim
return [{
"k": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
"v": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
"is_init": False,
} for _ in range(self.num_transformer_blocks)]
def _update_context_cache(self, block_latents: torch.Tensor, context: Any, camera_block: Any,
kv_cache: list[dict[str, Any]], crossattn_cache: list[dict[str, Any]], start: int,
frame_seq_length: int, target_dtype: torch.dtype, autocast_enabled: bool,
context_noise: float) -> None:
timestep = torch.full(
(block_latents.shape[0], block_latents.shape[2] * frame_seq_length),
self._context_noise_timestep(context_noise),
device=block_latents.device,
dtype=torch.long,
)
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
self.transformer(
hidden_states=block_latents.to(target_dtype),
encoder_hidden_states=torch.stack(context).to(target_dtype),
timestep=timestep,
y_camera=camera_block,
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=start * frame_seq_length,
)
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from dataclasses import dataclass
import numpy as np
import torch
from scipy.interpolate import interp1d
from scipy.spatial.transform import Rotation, Slerp
_ACTION_TO_MOTION = {
"w": "forward",
"a": "left",
"d": "right",
"s": "backward",
"j": "left_rot",
"l": "right_rot",
"i": "up_rot",
"k": "down_rot",
}
_TRANSLATION_BASE_UNIT = 1.0
_ROTATION_BASE_UNIT = 10.0
_INTRINSIC_ROW = [0.8, 0.5, 0.5, 0.5]
@dataclass
class DreamXCamera:
fx: float
fy: float
cx: float
cy: float
w2c_mat: np.ndarray
@property
def c2w_mat(self) -> np.ndarray:
return np.linalg.inv(self.w2c_mat)
@classmethod
def from_pose_row(cls, row: list[float]) -> DreamXCamera:
w2c_mat = np.eye(4, dtype=np.float64)
w2c_mat[:3, :] = np.asarray(row[7:], dtype=np.float64).reshape(3, 4)
return cls(
fx=float(row[1]),
fy=float(row[2]),
cx=float(row[3]),
cy=float(row[4]),
w2c_mat=w2c_mat,
)
def _translation_step(motion_type: str, current_pose: dict[str, np.ndarray], value: float, duration: int) -> np.ndarray:
if motion_type in ("forward", "backward"):
yaw = np.radians(current_pose["rotation"][1])
pitch = np.radians(current_pose["rotation"][0])
forward = np.array([-math.sin(yaw) * math.cos(pitch), math.sin(pitch), math.cos(yaw) * math.cos(pitch)])
direction = 1 if motion_type == "forward" else -1
return forward * value * direction / duration
if motion_type in ("left", "right"):
yaw = np.radians(current_pose["rotation"][1])
right = np.array([math.cos(yaw), 0.0, math.sin(yaw)])
direction = -1 if motion_type == "left" else 1
return right * value * direction / duration
return np.zeros(3)
def _rotation_step(motion_type: str, value: float, duration: int) -> np.ndarray:
if not motion_type.endswith("rot"):
return np.zeros(3)
axis = motion_type.split("_")[0]
rotation = np.zeros(3)
if axis == "left":
rotation[1] = value
elif axis == "right":
rotation[1] = -value
elif axis == "up":
rotation[0] = -value
elif axis == "down":
rotation[0] = value
return rotation / duration
def _euler_to_quaternion(angles: np.ndarray) -> list[float]:
pitch, yaw, roll = np.radians(angles)
cy = math.cos(yaw * 0.5)
sy = math.sin(yaw * 0.5)
cp = math.cos(pitch * 0.5)
sp = math.sin(pitch * 0.5)
cr = math.cos(roll * 0.5)
sr = math.sin(roll * 0.5)
return [
cy * cp * cr + sy * sp * sr,
cy * sp * cr + sy * cp * sr,
sy * cp * cr - cy * sp * sr,
cy * cp * sr - sy * sp * cr,
]
def _quaternion_to_rotation_matrix(quaternion: list[float]) -> np.ndarray:
qw, qx, qy, qz = quaternion
return np.array([
[1 - 2 * (qy**2 + qz**2), 2 * (qx * qy - qw * qz), 2 * (qx * qz + qw * qy)],
[2 * (qx * qy + qw * qz), 1 - 2 * (qx**2 + qz**2), 2 * (qy * qz - qw * qx)],
[2 * (qx * qz - qw * qy), 2 * (qy * qz + qw * qx), 1 - 2 * (qx**2 + qy**2)],
])
def _pose_rows_from_actions(action_seq: list[str], action_speed_list: list[float], duration: int) -> list[list[float]]:
if len(action_seq) != len(action_speed_list):
raise ValueError("action_seq and action_speed_list must have the same length")
positions: list[np.ndarray] = []
rotations: list[np.ndarray] = []
current_pose = {
"position": np.array([0.0, 0.0, 0.0]),
"rotation": np.array([0.0, 0.0, 0.0]),
}
for action_id, speed in zip(action_seq, action_speed_list, strict=True):
motion_types = [_ACTION_TO_MOTION[key] for key in list(action_id)]
translation_step = np.zeros(3)
rotation_step = np.zeros(3)
for motion_type in motion_types:
translation_step += _translation_step(motion_type, current_pose,
float(speed) * _TRANSLATION_BASE_UNIT, duration)
rotation_step += _rotation_step(motion_type, float(speed) * _ROTATION_BASE_UNIT, duration)
segment_positions = []
segment_rotations = []
for index in range(1, duration + 1):
segment_positions.append(current_pose["position"] + translation_step * index)
segment_rotations.append(current_pose["rotation"] + rotation_step * index)
current_pose["position"] = segment_positions[-1].copy()
current_pose["rotation"] = segment_rotations[-1].copy()
positions.extend(segment_positions)
rotations.extend(segment_rotations)
rows: list[list[float]] = [[0.0] + _INTRINSIC_ROW + [0.0, 0.0] +
[1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0]]
for index, (position, rotation) in enumerate(zip(positions, rotations, strict=False)):
rotation_matrix = _quaternion_to_rotation_matrix(_euler_to_quaternion(rotation))
translation = -rotation_matrix @ position
extrinsic = np.hstack([rotation_matrix, translation.reshape(3, 1)])
rows.append([float(index)] + _INTRINSIC_ROW + [0.0, 0.0] + extrinsic.flatten().tolist())
return rows
def _interpolate_camera_poses(
cameras: list[DreamXCamera],
src_indices: np.ndarray,
tgt_indices: np.ndarray,
) -> list[DreamXCamera]:
if len(cameras) <= 1:
return [cameras[0]] * len(tgt_indices) if cameras else []
src_rot_mat = np.array([camera.w2c_mat[:3, :3] for camera in cameras])
src_trans_vec = np.array([camera.w2c_mat[:3, 3] for camera in cameras])
dets = np.linalg.det(src_rot_mat)
flip_handedness = dets.size > 0 and np.median(dets) < 0.0
if flip_handedness:
flip_mat = np.diag([1.0, 1.0, -1.0]).astype(src_rot_mat.dtype)
src_rot_mat = src_rot_mat @ flip_mat
trans = interp1d(src_indices, src_trans_vec, axis=0, kind="linear", bounds_error=False,
fill_value="extrapolate")(tgt_indices)
quats = Rotation.from_matrix(src_rot_mat).as_quat().copy()
for index in range(1, len(quats)):
if np.dot(quats[index], quats[index - 1]) < 0:
quats[index] = -quats[index]
rot = Slerp(src_indices, Rotation.from_quat(quats))(tgt_indices).as_matrix()
if flip_handedness:
rot = rot @ flip_mat
ref = cameras[0]
result = []
for index in range(len(tgt_indices)):
w2c_mat = np.eye(4, dtype=np.float64)
w2c_mat[:3, :] = np.hstack([rot[index], trans[index].reshape(3, 1)])
result.append(DreamXCamera(ref.fx, ref.fy, ref.cx, ref.cy, w2c_mat))
return result
def _relative_c2w_poses(cameras: list[DreamXCamera]) -> np.ndarray:
abs_w2cs = [camera.w2c_mat for camera in cameras]
abs_c2ws = [camera.c2w_mat for camera in cameras]
target_cam_c2w = np.eye(4, dtype=np.float64)
abs2rel = target_cam_c2w @ abs_w2cs[0]
poses = [target_cam_c2w] + [abs2rel @ c2w for c2w in abs_c2ws[1:]]
return np.asarray(poses, dtype=np.float32)
def _invert_se3(transforms: torch.Tensor) -> torch.Tensor:
rotation_inv = transforms[..., :3, :3].transpose(-1, -2)
output = torch.zeros_like(transforms)
output[..., :3, :3] = rotation_inv
output[..., :3, 3] = -torch.einsum("...ij,...j->...i", rotation_inv, transforms[..., :3, 3])
output[..., 3, 3] = 1.0
return output
def build_dreamx_camera_condition(
action_seq: list[str],
action_speed_list: list[float],
*,
num_frames: int,
height: int,
width: int,
dtype: torch.dtype = torch.float32,
device: torch.device | str = "cpu",
) -> dict[str, torch.Tensor]:
del height, width # DreamX-World-5B-Cam uses fixed normalized intrinsics.
duration = math.ceil(num_frames / len(action_seq))
rows = _pose_rows_from_actions(action_seq, action_speed_list, duration)[:num_frames]
cameras = [DreamXCamera.from_pose_row(row) for row in rows]
latent_frame_count = 1 + (len(cameras) - 1) // 4
src_indices = np.arange(len(cameras), dtype=np.float64)
tgt_indices = np.linspace(0, len(cameras) - 1, latent_frame_count)
cameras = _interpolate_camera_poses(cameras, src_indices, tgt_indices)
c2ws = torch.as_tensor(_relative_c2w_poses(cameras), dtype=dtype, device=device)
viewmats = _invert_se3(c2ws)
intrinsics = torch.zeros((latent_frame_count, 3, 3), dtype=dtype, device=device)
intrinsics[:, 0, 0] = 969.6969696969696 / (960.0 * 2)
intrinsics[:, 1, 1] = 969.6969696969696 / (540.0 * 2)
intrinsics[:, 2, 2] = 1.0
return {"viewmats": viewmats, "K": intrinsics}
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
"""Compatibility exports for DreamX-World pipeline configs."""
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BARPipelineConfig,
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_ar_dit_config,
make_dreamx_world_5b_cam_dit_config,
make_dreamx_world_5b_cam_text_encoder_config,
make_dreamx_world_5b_cam_vae_config,
)
__all__ = [
"DreamXWorld5BARPipelineConfig",
"DreamXWorld5BCamPipelineConfig",
"make_dreamx_world_5b_ar_dit_config",
"make_dreamx_world_5b_cam_dit_config",
"make_dreamx_world_5b_cam_text_encoder_config",
"make_dreamx_world_5b_cam_vae_config",
]
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive pipeline entrypoint."""
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
from fastvideo.pipelines.basic.dreamx_world.stages import (
DreamXWorldCameraConditioningStage,
DreamXWorldImageVAEEncodingStage,
)
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
logger = init_logger(__name__)
class DreamXWorldARPipeline(LoRAPipeline, ComposedPipelineBase):
"""DreamX-World-5B autoregressive causal camera pipeline."""
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
pipeline_config_cls = DreamXWorld5BARPipelineConfig
sampling_params_cls = SamplingParam
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
self.modules["scheduler"].set_timesteps(1000)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=DreamXWorldImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
self.add_stage(stage_name="denoising_stage",
stage=DreamXWorldARCausalDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
logger.info("DreamXWorldARPipeline initialized with autoregressive causal denoising")
EntryClass = DreamXWorldARPipeline
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World video pipeline entrypoint."""
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
from fastvideo.pipelines.basic.dreamx_world.stages import DreamXWorldCameraConditioningStage
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
logger = init_logger(__name__)
class DreamXWorldPipeline(LoRAPipeline, ComposedPipelineBase):
"""DreamX-World-5B-Cam pipeline with native FastVideo camera conditioning."""
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
pipeline_config_cls = DreamXWorld5BCamPipelineConfig
sampling_params_cls = SamplingParam
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(
stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
),
)
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
self.add_stage(
stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self,
),
)
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
logger.info("DreamXWorldPipeline initialized with native camera conditioning")
EntryClass = DreamXWorldPipeline
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_NEGATIVE_PROMPT_CN = ("色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,"
"静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,"
"多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,"
"形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,"
"背景人很多,倒着走")
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="DreamX-World camera-conditioned denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
DREAMX_WORLD_5B_CAM = InferencePreset(
name="dreamx_world_5b_cam",
version=1,
model_family="dreamx_world",
description="DreamX-World 5B camera-control video generation",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 480,
"width": 832,
"num_frames": 161,
"fps": 16,
"guidance_scale": 5.0,
"num_inference_steps": 30,
"negative_prompt": _NEGATIVE_PROMPT_CN,
},
)
DREAMX_WORLD_5B_AR = InferencePreset(
name="dreamx_world_5b_ar",
version=1,
model_family="dreamx_world",
description="DreamX-World 5B autoregressive camera-control generation",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 704,
"width": 1280,
"num_frames": 1005,
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": _NEGATIVE_PROMPT_CN,
},
)
ALL_PRESETS = (DREAMX_WORLD_5B_CAM, DREAMX_WORLD_5B_AR)
@@ -0,0 +1,156 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline stages."""
from __future__ import annotations
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
build_dreamx_camera_condition, )
DREAMX_Y_CAMERA_KEY = "dreamx_y_camera"
logger = init_logger(__name__)
class DreamXWorldCameraConditioningStage(PipelineStage):
"""Build PRoPE camera conditioning for DreamX-World-5B-Cam."""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
del fastvideo_args
if DREAMX_Y_CAMERA_KEY in batch.extra:
return batch
action_seq = batch.extra.get("dreamx_action_seq", batch.action_list)
action_speed_list = batch.extra.get("dreamx_action_speed_list", batch.action_speed_list)
if action_seq is None:
action_seq = ["w"]
if action_speed_list is None:
action_speed_list = [4]
if isinstance(action_seq, str):
action_seq = [action_seq]
if isinstance(action_speed_list, int | float):
action_speed_list = [action_speed_list]
if len(action_speed_list) == 1 and len(action_seq) > 1:
action_speed_list = list(action_speed_list) * len(action_seq)
action_speed_list = [float(speed) for speed in action_speed_list]
height = int(batch.height) if batch.height is not None else 704
width = int(batch.width) if batch.width is not None else 1280
num_frames = int(batch.num_frames)
dtype = batch.latents.dtype if torch.is_tensor(batch.latents) else torch.float32
device = batch.latents.device if torch.is_tensor(batch.latents) else "cpu"
y_camera = build_dreamx_camera_condition(
list(action_seq),
action_speed_list,
num_frames=num_frames,
height=height,
width=width,
dtype=dtype,
device=device,
)
batch.extra[DREAMX_Y_CAMERA_KEY] = {key: value.unsqueeze(0) for key, value in y_camera.items()}
return batch
def verify_output(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> VerificationResult:
del fastvideo_args
result = VerificationResult()
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY)
result.add_check("dreamx_y_camera", y_camera, lambda value: isinstance(value, dict))
if isinstance(y_camera, dict):
result.add_check("dreamx_y_camera.viewmats", y_camera.get("viewmats"), torch.is_tensor)
result.add_check("dreamx_y_camera.K", y_camera.get("K"), torch.is_tensor)
return result
class DreamXWorldImageVAEEncodingStage(PipelineStage):
"""Encode the conditioning image into the first-frame latent.
Official AR-forcing flow (AMAP-ML/DreamX-World inference_ar_forcing.py):
the input image is resized, normalized to [-1, 1], VAE-encoded
deterministically, and written into frame 0 of the noise — the causal
denoiser then treats frame 0 as clean context. This stage produces
``batch.image_latent`` ([B, C, 1, H_lat, W_lat]); the injection into
the latents happens in DreamXWorldARCausalDenoisingStage.
"""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
if batch.pil_image is None:
# No conditioning image: the causal denoiser falls back to
# running from pure noise (frame 0 uninitialized). Warn loudly —
# this pipeline is registered I2V and the official flow always
# forces from a frame.
logger.warning("DreamXWorldARPipeline called without an input image; "
"first-frame context will be noise (T2V-style). Pass an "
"image for the official AR-forcing behavior.")
return batch
from fastvideo.platforms import get_local_torch_device
from fastvideo.utils import PRECISION_TO_TYPE
device = get_local_torch_device()
image = batch.pil_image
if not isinstance(image, torch.Tensor):
import numpy as np
import PIL.Image
assert isinstance(image, PIL.Image.Image)
width = batch.width if isinstance(batch.width, int) else batch.width[0]
height = batch.height if isinstance(batch.height, int) else batch.height[0]
image = image.convert("RGB").resize((width, height), PIL.Image.Resampling.LANCZOS)
arr = torch.from_numpy(np.asarray(image)).float().permute(2, 0, 1) / 255.0
image = (arr - 0.5) / 0.5 # official Normalize([0.5], [0.5])
image = image.unsqueeze(0) # [1, C, H, W]
if image.dim() == 4:
image = image.unsqueeze(2) # [B, C, 1, H, W]
elif image.dim() == 5:
image = image[:, :, :1]
image = image.to(device=device, dtype=torch.float32)
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
self.vae = self.vae.to(device)
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if not vae_autocast_enabled:
image = image.to(vae_dtype)
encoder_output = self.vae.encode(image)
# Official encode_to_latent is deterministic ((mean - mu) / sigma per
# channel); the posterior mean + shift/scale is the FastVideo
# equivalent of that normalization.
latent = encoder_output.mean
if getattr(self.vae, "shift_factor", None) is not None:
shift = self.vae.shift_factor
latent = latent - (shift.to(latent.device, latent.dtype) if isinstance(shift, torch.Tensor) else shift)
scale = self.vae.scaling_factor
latent = latent * (scale.to(latent.device, latent.dtype) if isinstance(scale, torch.Tensor) else scale)
batch.image_latent = latent
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
return result
+4
View File
@@ -34,6 +34,7 @@ class PipelineStage(ABC):
composed with other stages to create a complete pipeline. Each stage is responsible
for a specific part of the process, such as prompt encoding, latent preparation, etc.
"""
performance_component_metric: str | None = None
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
"""
@@ -155,6 +156,9 @@ class PipelineStage(ABC):
logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000)
batch.logging_info.add_stage_execution_time(stage_key, execution_time)
batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name)
component_metric = self.performance_component_metric
if component_metric is not None:
batch.logging_info.add_stage_metric(stage_key, "component_metric", component_metric)
except Exception as e:
torch.cuda.synchronize()
execution_time = time.perf_counter() - start_time
+1
View File
@@ -28,6 +28,7 @@ class DecodingStage(PipelineStage):
This stage handles the decoding of latent representations into the final
output format (e.g., pixel values).
"""
performance_component_metric = "vae_decode_time_s"
def __init__(self, vae, pipeline=None) -> None:
self.vae: ParallelTiledVAE = vae
+21
View File
@@ -51,6 +51,7 @@ class DenoisingStage(PipelineStage):
This stage handles the iterative denoising process that transforms
the initial noise into the final output.
"""
performance_component_metric = "dit_time_s"
def __init__(self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None) -> None:
super().__init__()
@@ -190,6 +191,19 @@ class DenoisingStage(PipelineStage):
},
)
dreamx_y_camera = batch.extra.get("dreamx_y_camera", batch.extra.get("y_camera"))
if isinstance(dreamx_y_camera, dict):
dreamx_y_camera = {
key: value.to(device=local_device, dtype=target_dtype) if torch.is_tensor(value) else value
for key, value in dreamx_y_camera.items()
}
dreamx_camera_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"y_camera": dreamx_y_camera,
},
)
for key in ("flux2_txt_ids", "flux2_img_ids"):
value = batch.extra.get(key)
if torch.is_tensor(value):
@@ -241,7 +255,11 @@ class DenoisingStage(PipelineStage):
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
vae_device = next(self.vae.parameters()).device
self.vae = self.vae.to(local_device)
z = self.vae.encode(batch.pil_image).mean.float()
if getattr(fastvideo_args, "vae_cpu_offload", False):
self.vae = self.vae.to(vae_device)
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
z -= self.vae.shift_factor.to(z.device, z.dtype)
@@ -494,6 +512,7 @@ class DenoisingStage(PipelineStage):
**pos_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
@@ -536,6 +555,7 @@ class DenoisingStage(PipelineStage):
**neg_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
@@ -1187,6 +1207,7 @@ class Cosmos25V2WDenoisingStage(Cosmos25DenoisingStage):
class Cosmos25AutoDenoisingStage(PipelineStage):
"""Route Cosmos 2.5 denoising to T2W vs V2W/I2W."""
performance_component_metric = "dit_time_s"
def __init__(self, transformer, scheduler) -> None:
super().__init__()
@@ -24,6 +24,7 @@ class TextEncodingStage(PipelineStage):
This stage handles the encoding of text prompts into the embedding space
expected by the diffusion model.
"""
performance_component_metric = "text_encoder_time_s"
def __init__(self, text_encoders, tokenizers) -> None:
"""
@@ -350,6 +351,7 @@ class Cosmos25TextEncodingStage(PipelineStage):
Cosmos 2.5 uses Reason1 (Qwen2.5-VL) and relies on the encoder's
`compute_text_embeddings_online()`.
"""
performance_component_metric = "text_encoder_time_s"
def __init__(self, text_encoder) -> None:
super().__init__()
+40
View File
@@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
Cosmos25Config,
Cosmos25_14BConfig,
)
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
@@ -773,6 +774,42 @@ def _register_configs() -> None:
model_family="wan",
default_preset="wan_2_2_ti2v_5b",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=DreamXWorld5BCamPipelineConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=[
"FastVideo/DreamX-World-5B-Cam-Diffusers",
],
model_detectors=[
# Pattern also catches the raw GD-ML/DreamX-World-5B-Cam id and
# local converted dirs. Mutually exclusive with the AR detector
# below: Cam requires an explicit "cam" marker so hyphenated AR
# local paths (e.g. /ckpts/dreamx-world-5b-converted) don't
# first-match here — detector resolution is first-match in
# registration order.
lambda path:
("dreamx-world" in path.lower() and "cam" in path.lower()) or "dreamxworldpipeline" in path.lower()
],
model_family="dreamx_world",
default_preset="dreamx_world_5b_cam",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=DreamXWorld5BARPipelineConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=[
"FastVideo/DreamX-World-5B-Diffusers",
],
model_detectors=[
lambda path:
("dreamx-world-5b" in path.lower() and "cam" not in path.lower()) or "dreamxworldarpipeline" in path.lower(
)
],
model_family="dreamx_world",
default_preset="dreamx_world_5b_ar",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
@@ -951,6 +988,8 @@ def _register_presets() -> None:
from fastvideo.api.presets import register_preset
from fastvideo.pipelines.basic.cosmos.presets import (
ALL_PRESETS as COSMOS_PRESETS, )
from fastvideo.pipelines.basic.dreamx_world.presets import (
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
from fastvideo.pipelines.basic.gamecraft.presets import (
ALL_PRESETS as GAMECRAFT_PRESETS, )
from fastvideo.pipelines.basic.gen3c.presets import (
@@ -984,6 +1023,7 @@ def _register_presets() -> None:
all_preset_groups = (
COSMOS_PRESETS,
DREAMX_WORLD_PRESETS,
FLUX2_PRESETS,
GAMECRAFT_PRESETS,
GEN3C_PRESETS,
@@ -3,15 +3,16 @@
Landed in PR #1225 slice 5 (Attn-QAT 5/12). The resolver centralises the
varlen-flash-attn import-fallback logic that several backends
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The fallback
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The resolution
order is:
1. ``fastvideo.attention.utils.flash_attn_cute``
1. ``fastvideo.attention.utils.flash_attn_cute`` -- only when
``FASTVIDEO_FA4=1`` (explicit opt-in), and then it must import or the
resolver raises RuntimeError instead of falling through
2. ``flash_attn_interface``
3. ``flash_attn``
These tests verify that the resolver picks the highest-priority impl
available and falls through cleanly on ``ImportError``. CPU-only, no
These tests verify the opt-in gate and the FA3/FA2 fallthrough. CPU-only, no
flash-attn install required.
"""
@@ -35,8 +36,36 @@ def _reload_resolver_module():
return importlib.import_module("fastvideo.attention.utils.flash_attn_no_pad")
def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
"""When ``flash_attn_cute`` is unimportable, resolver tries the next impl."""
def test_resolver_skips_cute_without_opt_in(monkeypatch) -> None:
"""Without ``FASTVIDEO_FA4=1`` the resolver must not even attempt the cute
import."""
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
attempted: list[str] = []
real_import = builtins.__import__
def spying_import(name, globals=None, locals=None, fromlist=(), level=0):
attempted.append(name)
return real_import(name, globals, locals, fromlist, level)
monkeypatch.setattr(builtins, "__import__", spying_import)
mod = _reload_resolver_module()
resolved = mod._resolve_flash_attn_varlen_func()
assert resolved is not None
assert resolved.__name__ == "flash_attn_varlen_func"
assert "fastvideo.attention.utils.flash_attn_cute" not in attempted
def test_resolver_raises_when_opted_in_but_cute_unavailable(monkeypatch) -> None:
"""With ``FASTVIDEO_FA4=1`` an unimportable cute build fails loudly instead
of silently falling through to FA3/FA2.
The resolver runs at module import time, so the reload itself must raise.
It raises RuntimeError (not ImportError) so importers that treat
ImportError as "flash-attn not installed" (``bsa_attn.py``) cannot swallow
the opted-in failure.
"""
monkeypatch.setenv("FASTVIDEO_FA4", "1")
real_import = builtins.__import__
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
@@ -46,21 +75,17 @@ def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
monkeypatch.setattr(builtins, "__import__", patched_import)
mod = _reload_resolver_module()
resolved = mod._resolve_flash_attn_varlen_func()
assert resolved is not None
assert resolved.__name__ == "flash_attn_varlen_func"
with pytest.raises(RuntimeError, match="cute disabled for test"):
_reload_resolver_module()
def test_resolver_returns_flash_attn_when_cute_and_interface_unavailable(monkeypatch) -> None:
def test_resolver_returns_flash_attn_when_interface_unavailable(monkeypatch) -> None:
"""The terminal fallback is the plain ``flash_attn`` import."""
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
real_import = builtins.__import__
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
if name in {
"fastvideo.attention.utils.flash_attn_cute",
"flash_attn_interface",
}:
if name == "flash_attn_interface":
raise ImportError(f"{name} disabled for test")
return real_import(name, globals, locals, fromlist, level)
@@ -0,0 +1,85 @@
# SPDX-License-Identifier: Apache-2.0
"""Guard: every test directory must be collected by some CI lane or be on
the explicit allowlist below.
Three separate incidents on 2026-07-05 found test files that no CI lane
ever collects (fastvideo/tests/stages/, tests/local_tests/ additions in
PR #1509, and this sweep found seven dark directories in total): the tests
pass review, merge, and then silently never run. This test makes going
dark an explicit, reviewed decision instead of an accident: adding a new
test directory fails CI until it is either wired into a lane or
allowlisted here with a reason.
Pure text analysis — no fastvideo imports, no GPU, no torch.
"""
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[3]
TESTS_ROOT = REPO_ROOT / "fastvideo" / "tests"
# Files whose text constitutes "a CI lane references this directory".
CI_SOURCES = [
TESTS_ROOT / "modal" / "pr_test.py",
TESTS_ROOT / "modal" / "ssim_test.py",
*sorted((REPO_ROOT / ".buildkite").rglob("*.yml")),
*sorted((REPO_ROOT / ".buildkite").rglob("*.sh")),
]
# Directories that intentionally have no CI lane today. Every entry needs a
# reason; remove the entry when the directory gets wired into a lane.
# State as found on 2026-07-05 — these SHOULD shrink over time, not grow.
ALLOWLIST = {
"attention": "no lane yet — GPU attention-backend tests, run manually",
"audio": "no lane yet — audio encoder tests, run manually",
"distributed": "no lane yet — multi-GPU torchrun tests, run manually",
"hooks": "no lane yet — run manually",
"layers": "no lane yet — torchrun FSDP dispatch tests, run manually",
"nightly": "by design: nightly cadence, not per-PR",
"ops": "no lane yet — GPU op tests, run manually",
"modal": "CI infrastructure itself, not a test suite",
}
def _dirs_with_tests() -> list[str]:
dirs = []
for child in sorted(TESTS_ROOT.iterdir()):
if child.is_dir() and any(child.rglob("test_*.py")):
dirs.append(child.name)
return dirs
def _ci_text() -> str:
return "\n".join(
src.read_text(errors="replace") for src in CI_SOURCES if src.exists())
def test_every_test_directory_is_collected_or_allowlisted():
ci_text = _ci_text()
dark = [
name for name in _dirs_with_tests()
if f"tests/{name}" not in ci_text and name not in ALLOWLIST
]
assert not dark, (
f"Test directories not referenced by any CI lane and not "
f"allowlisted: {dark}. Wire them into a lane in "
f"fastvideo/tests/modal/pr_test.py (or a Buildkite step), or add an "
f"allowlist entry with a reason in {__file__}.")
def test_local_tests_stays_out_of_ci():
# tests/local_tests/ (repo root) is developer-local by design (author
# decision, 2026-07-05): parity scaffolds and machine-specific checks
# that must never gate CI. Fail if any CI source starts collecting it.
assert "tests/local_tests" not in _ci_text(), (
"tests/local_tests/ is local-only by design; remove the CI "
"reference or move the tests into a fastvideo/tests/ lane.")
def test_allowlist_entries_are_still_real_directories():
# A stale allowlist hides regressions; entries must track reality.
missing = [
name for name in ALLOWLIST
if name != "modal" and not (TESTS_ROOT / name).is_dir()
]
assert not missing, (
f"Allowlisted directories no longer exist — remove them: {missing}")
@@ -0,0 +1,203 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import json
import os
import subprocess
from pathlib import Path
from typing import Any
import pytest
import torch
import torch.distributed as dist
from torch.distributed import init_device_mesh
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed.tensor import DTensor
from fastvideo.layers.layernorm import RMSNorm
WORLD_SIZE = 2
HIDDEN_SIZE = 8
SEED = 1379
REPO_ROOT = Path(__file__).resolve().parents[3]
def _run_torchrun(script_path: Path, mode: str, output_path: Path) -> None:
# --standalone binds the rendezvous port atomically, avoiding the
# free-port-probe race a hand-picked --master_port would have.
cmd = [
"torchrun",
"--standalone",
"--nproc_per_node",
str(WORLD_SIZE),
str(script_path),
"--rmsnorm-fsdp-worker",
"--mode",
mode,
"--output",
str(output_path),
]
env = os.environ.copy()
env.setdefault("TORCHDYNAMO_DISABLE", "1")
try:
process = subprocess.run(
cmd,
capture_output=True,
text=True,
env=env,
timeout=120,
)
except subprocess.TimeoutExpired as error:
raise RuntimeError(
f"{mode} worker timed out after 120 seconds\n"
f"STDOUT:\n{error.stdout}\n"
f"STDERR:\n{error.stderr}"
) from error
if process.returncode != 0:
raise RuntimeError(
f"{mode} worker failed with code {process.returncode}\n"
f"STDOUT:\n{process.stdout}\n"
f"STDERR:\n{process.stderr}"
)
def _summarize_tensor(tensor: torch.Tensor | Any) -> dict[str, Any]:
return {
"type": type(tensor).__name__,
"is_dtensor": isinstance(tensor, DTensor),
"shape": list(tensor.shape) if hasattr(tensor, "shape") else None,
"device": str(tensor.device) if hasattr(tensor, "device") else None,
"dtype": str(tensor.dtype) if hasattr(tensor, "dtype") else None,
}
def _run_worker(mode: str, output_path: Path) -> None:
if mode not in {
"module_no_offload",
"direct_no_offload",
"module_cpu_offload",
"direct_cpu_offload",
}:
raise ValueError(f"Unsupported mode: {mode}")
dist.init_process_group("nccl")
rank = dist.get_rank()
world_size = dist.get_world_size()
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
torch.manual_seed(SEED + rank)
try:
mesh = init_device_mesh("cuda", (world_size,))
norm = RMSNorm(HIDDEN_SIZE, eps=1e-6, has_weight=True).to(device)
with torch.no_grad():
norm.weight.fill_(1.0)
fsdp_kwargs: dict[str, Any] = {"mesh": mesh}
if mode.endswith("cpu_offload"):
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=False)
# fully_shard is applied to the bare RMSNorm to make the hook bypass
# observable. Production sharding (fsdp_load.shard_model) only wraps
# whole transformer blocks, whose pre-forward all-gather localizes norm
# weights before the qk-norm call sites run, so this pins the dispatch
# invariant rather than reproducing a production topology.
fully_shard(norm, **fsdp_kwargs)
x = torch.randn(2, 3, HIDDEN_SIZE, device=device, dtype=torch.bfloat16)
call_kind = "direct" if mode.startswith("direct") else "module"
try:
if call_kind == "direct":
output = norm.forward_native(x)
else:
output = norm(x)
torch.cuda.synchronize(device)
result = {
"rank": rank,
"ok": True,
"mode": mode,
"weight": _summarize_tensor(norm.weight),
"output": _summarize_tensor(output),
}
except Exception as exc:
result = {
"rank": rank,
"ok": False,
"mode": mode,
"error_type": type(exc).__name__,
"error": str(exc),
"weight": _summarize_tensor(norm.weight),
}
gathered = [None for _ in range(world_size)] if rank == 0 else None
dist.gather_object(result, object_gather_list=gathered, dst=0)
if rank == 0:
output_path.write_text(json.dumps(gathered, indent=2), encoding="utf-8")
dist.barrier()
finally:
dist.destroy_process_group()
@pytest.mark.parametrize(
("mode", "expect_ok"),
[
("module_no_offload", True),
("direct_no_offload", False),
("module_cpu_offload", True),
("direct_cpu_offload", False),
],
)
def test_rmsnorm_forward_native_bypasses_fsdp_hooks(mode: str, expect_ok: bool, tmp_path: Path) -> None:
if not torch.cuda.is_available():
pytest.skip("This test requires CUDA.")
if torch.cuda.device_count() < WORLD_SIZE:
pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.")
output_path = tmp_path / f"{mode}.json"
_run_torchrun(Path(__file__).resolve(), mode, output_path)
results = json.loads(output_path.read_text(encoding="utf-8"))
print(f"\n{mode} results:\n{json.dumps(results, indent=2)}")
if expect_ok:
failures = [result for result in results if not result["ok"]]
assert not failures, json.dumps(results, indent=2)
return
successes = [result for result in results if result["ok"]]
assert not successes, json.dumps(results, indent=2)
error_text = "\n".join(result.get("error", "") for result in results)
# Pin the specific bypassed-hook failure: "got mixed torch.Tensor and
# DTensor" ("Tensor" alone is a substring of "DTensor", so it adds nothing).
assert "mixed" in error_text and "DTensor" in error_text, json.dumps(results, indent=2)
def test_no_direct_forward_native_calls_in_models() -> None:
"""Direct .forward_native(...) calls bypass nn.Module.__call__ and FSDP
hooks (issue #1379); model code must use module dispatch instead."""
models_dir = REPO_ROOT / "fastvideo" / "models"
offenders = [
str(path.relative_to(REPO_ROOT))
for path in sorted(models_dir.rglob("*.py"))
if ".forward_native(" in path.read_text(encoding="utf-8")
]
assert not offenders, f"Replace .forward_native(...) with module dispatch in: {offenders}"
def _parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--rmsnorm-fsdp-worker", action="store_true")
parser.add_argument("--mode", type=str, default=None)
parser.add_argument("--output", type=str, default=None)
return parser.parse_args()
if __name__ == "__main__":
args = _parse_args()
if not args.rmsnorm_fsdp_worker:
raise SystemExit("This module is intended to be run by pytest.")
if args.mode is None or args.output is None:
raise SystemExit("--mode and --output are required in worker mode.")
_run_worker(mode=args.mode, output_path=Path(args.output))
+9 -7
View File
@@ -32,7 +32,8 @@ import modal
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
try:
from modal_image_utils import resolve_image_ref # noqa: E402
from modal_image_utils import ( # noqa: E402
resolve_image_ref, resolve_uv_torch_backend)
except ModuleNotFoundError:
# Remote Modal containers re-import this module but mount only the
# entrypoint file; the digest resolution already happened at local
@@ -40,6 +41,9 @@ except ModuleNotFoundError:
def resolve_image_ref(image_ref: str) -> str:
return image_ref
def resolve_uv_torch_backend(image_tag: str) -> str | None:
return os.environ.get("UV_TORCH_BACKEND")
app = modal.App("fastvideo-gpu-job")
REPO_DIR = "/FastVideo"
@@ -72,12 +76,7 @@ local_secrets = modal.Secret.from_dict({
# Mutable tags inherit the registry image's baked backend, including custom
# FASTVIDEO_MODAL_IMAGE overrides. Explicit CUDA tags also work with older
# images that predate the baked setting, and a caller override always wins.
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
if not uv_torch_backend_override:
if "cuda13" in IMAGE_TAG.lower():
uv_torch_backend_override = "cu130"
elif "cuda12.6" in IMAGE_TAG.lower():
uv_torch_backend_override = "cu126"
uv_torch_backend_override = resolve_uv_torch_backend(IMAGE_TAG)
image = (
modal.Image.from_registry(IMAGE_REF, add_python="3.12")
@@ -98,6 +97,9 @@ image = (
"TOKENIZERS_PARALLELISM": "false",
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
# references. Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
})
)
@@ -5,6 +5,7 @@ through the ``fastvideo`` package, so they keep working on thin CI hosts that
have ``modal`` but not torch.
"""
import json
import os
import urllib.request
_REGISTRY = "ghcr.io"
@@ -55,3 +56,23 @@ def resolve_image_ref(image_ref: str) -> str:
print(f"WARNING: could not resolve {image_ref} to a digest ({error}); "
"Modal may reuse a stale cached image for this tag.")
return image_ref
def resolve_uv_torch_backend(image_tag: str) -> str | None:
"""UV_TORCH_BACKEND for a launcher image tag.
A caller-set UV_TORCH_BACKEND always wins. Otherwise sniff the CUDA
version from an explicit image tag (cuda13 -> cu130, cuda12.6 -> cu126)
so uv resolves torch against the image's toolkit. Mutable tags (e.g.
py3.12-latest) return None and inherit the registry image's baked
backend, which keeps a latest-tag CUDA transition safe.
"""
override = os.environ.get("UV_TORCH_BACKEND")
if override:
return override
tag = image_tag.lower()
if "cuda13" in tag:
return "cu130"
if "cuda12.6" in tag:
return "cu126"
return None
+11 -8
View File
@@ -5,7 +5,8 @@ import modal
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
try:
from modal_image_utils import resolve_image_ref # noqa: E402
from modal_image_utils import ( # noqa: E402
resolve_image_ref, resolve_uv_torch_backend)
except ModuleNotFoundError:
# Remote Modal containers re-import this module but mount only the
# entrypoint file; the digest resolution already happened at local
@@ -13,6 +14,9 @@ except ModuleNotFoundError:
def resolve_image_ref(image_ref: str) -> str:
return image_ref
def resolve_uv_torch_backend(image_tag: str) -> str | None:
return os.environ.get("UV_TORCH_BACKEND")
app = modal.App()
model_vol = modal.Volume.from_name("hf-model-weights")
@@ -24,12 +28,7 @@ print(f"Using image: {image_ref}")
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
# transition safe. Explicit CUDA tags also work with older images that predate
# the baked setting, and a caller override always wins.
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
if not uv_torch_backend_override:
if "cuda13" in image_tag.lower():
uv_torch_backend_override = "cu130"
elif "cuda12.6" in image_tag.lower():
uv_torch_backend_override = "cu126"
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
image = (modal.Image.from_registry(
image_ref, add_python="3.12"
@@ -61,6 +60,10 @@ image = (modal.Image.from_registry(
**({
"UV_TORCH_BACKEND": uv_torch_backend_override
} if uv_torch_backend_override else {}),
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
# SSIM/perf baselines. Caller override wins.
"FASTVIDEO_FA4":
os.environ.get("FASTVIDEO_FA4", "1"),
"HF_REPO_ID":
"FastVideo/performance-tracking",
}))
@@ -273,7 +276,7 @@ def run_self_forcing_tests():
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test(
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
)
+9 -7
View File
@@ -13,7 +13,8 @@ import modal
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
try:
from modal_image_utils import resolve_image_ref # noqa: E402
from modal_image_utils import ( # noqa: E402
resolve_image_ref, resolve_uv_torch_backend)
except ModuleNotFoundError:
# Remote Modal containers re-import this module but mount only the
# entrypoint file; the digest resolution already happened at local
@@ -21,6 +22,9 @@ except ModuleNotFoundError:
def resolve_image_ref(image_ref: str) -> str:
return image_ref
def resolve_uv_torch_backend(image_tag: str) -> str | None:
return os.environ.get("UV_TORCH_BACKEND")
app = modal.App()
model_vol = modal.Volume.from_name("hf-model-weights")
@@ -32,12 +36,7 @@ print(f"Using image: {image_ref}")
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
# transition safe. Explicit CUDA tags also work with older images that predate
# the baked setting, and a caller override always wins.
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
if not uv_torch_backend_override:
if "cuda13" in image_tag.lower():
uv_torch_backend_override = "cu130"
elif "cuda12.6" in image_tag.lower():
uv_torch_backend_override = "cu126"
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
image = (
modal.Image.from_registry(image_ref, add_python="3.12")
@@ -64,6 +63,9 @@ image = (
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"IMAGE_VERSION": image_version,
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
# FA4 is opt-in (FASTVIDEO_FA4); the SSIM references were seeded
# with FA4 inference, so keep it enabled in CI. Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
}
)
)
+92 -75
View File
@@ -8,8 +8,8 @@ This script:
baseline-eligible successful records (filtered by gpu_type),
4) writes normalized records back to the HF dataset repo according to
PERF_UPLOAD_POLICY,
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
(default 5%).
5) exits non-zero if any gated metric exceeds both its percent and absolute
regression floors.
"""
import glob
@@ -21,21 +21,36 @@ from datetime import datetime, timezone
from typing import Any
try:
from .hf_store import (
from fastvideo.performance.hf_store import (
load_records_for_model,
safe_float,
sanitize,
sync_from_hf,
upload_record,
)
from fastvideo.performance.metric_policy import (
MetricPolicy,
regression_delta,
resolve_metric_policies,
serialize_metric_thresholds,
)
except ImportError:
from hf_store import (
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
if repo_root not in sys.path:
sys.path.insert(0, repo_root)
from fastvideo.performance.hf_store import (
load_records_for_model,
safe_float,
sanitize,
sync_from_hf,
upload_record,
)
from fastvideo.performance.metric_policy import (
MetricPolicy,
regression_delta,
resolve_metric_policies,
serialize_metric_thresholds,
)
RESULTS_DIR = os.path.join(
os.path.dirname(os.path.abspath(__file__)),
@@ -46,25 +61,9 @@ TRACKING_ROOT = os.environ.get(
"/tmp/perf-tracking",
)
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
UPLOAD_POLICY = os.environ.get("PERF_UPLOAD_POLICY", "never").strip().lower()
VALID_UPLOAD_POLICIES = {"never", "pass", "always"}
VALID_RUN_SOURCES = {"pr", "local", "scheduled_main", "unknown"}
METRICS = (
("latency", "Latency", 3),
("throughput", "Throughput", 3),
("memory", "Memory", 1),
("text_encoder_time_s", "Text Enc", 3),
("dit_time_s", "DiT", 3),
("vae_decode_time_s", "VAE Decode", 3),
)
LOWER_IS_BETTER_METRICS = {
"latency",
"memory",
"text_encoder_time_s",
"dit_time_s",
"vae_decode_time_s",
}
def _should_persist_tracking() -> bool:
@@ -168,6 +167,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
text_encoder_time = safe_float(result.get("text_encoder_time_s"))
dit_time = safe_float(result.get("dit_time_s"))
vae_decode_time = safe_float(result.get("vae_decode_time_s"))
metric_policies = resolve_metric_policies(result.get("regression_thresholds"))
return {
"model_id": model_id,
@@ -180,6 +180,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
"text_encoder_time_s": text_encoder_time,
"dit_time_s": dit_time,
"vae_decode_time_s": vae_decode_time,
"regression_thresholds": serialize_metric_thresholds(metric_policies),
"success": True,
**_record_metadata(_detect_run_source(), result),
}
@@ -229,55 +230,41 @@ def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
return statistics.median(values)
def _metric_policy_summary(policy: MetricPolicy) -> str:
gated = "gated" if policy.gated else "info"
return (
f"{gated}, >{policy.threshold_percent * 100:.1f}% "
f"and >{policy.threshold_absolute:.{policy.precision}f}"
)
def _check_regressions(
current: dict[str, Any],
baseline_records: list[dict[str, Any]],
max_regression: float,
metric_policies: tuple[MetricPolicy, ...],
) -> list[str]:
failures: list[str] = []
for metric, _label, _precision in METRICS:
if metric not in LOWER_IS_BETTER_METRICS:
for policy in metric_policies:
baseline = _baseline_metric(baseline_records, policy.key)
curr = safe_float(current.get(policy.key))
if baseline is None or curr is None:
continue
baseline = _baseline_metric(baseline_records, metric)
curr = safe_float(current.get(metric))
if baseline is None or curr is None or baseline <= 0:
delta = regression_delta(policy, curr, baseline)
if delta is None or not delta.regressed:
continue
regression = (curr - baseline) / baseline
if regression > max_regression:
failures.append(f"{current['model_id']} {metric} regressed by "
f"{regression * 100:.1f}% "
f"(current={curr:.3f}, baseline_median={baseline:.3f})")
baseline_tp = _baseline_metric(baseline_records, "throughput")
curr_tp = safe_float(current.get("throughput"))
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
regression = (baseline_tp - curr_tp) / baseline_tp
if regression > max_regression:
failures.append(f"{current['model_id']} throughput regressed by "
f"{regression * 100:.1f}% "
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})")
failures.append(
f"{current['model_id']} {policy.key} regressed by "
f"{delta.percent * 100:.1f}% and "
f"{delta.absolute:.{policy.precision}f} "
f"(current={curr:.{policy.precision}f}, "
f"baseline_median={baseline:.{policy.precision}f}, "
f"threshold={_metric_policy_summary(policy)})"
)
return failures
def _metric_delta_percent(
metric: str,
current: dict[str, Any],
baseline_records: list[dict[str, Any]],
) -> float | None:
curr = safe_float(current.get(metric))
baseline = _baseline_metric(baseline_records, metric)
if curr is None or baseline is None or baseline <= 0:
return None
if metric in LOWER_IS_BETTER_METRICS:
return (curr - baseline) / baseline * 100.0
if metric == "throughput":
return (baseline - curr) / baseline * 100.0
return None
def _compact_value(value: float | None, precision: int = 3) -> str:
if value is None:
return "n/a"
@@ -287,23 +274,42 @@ def _compact_value(value: float | None, precision: int = 3) -> str:
def _build_summary_row(
record: dict[str, Any],
baseline_records: list[dict[str, Any]],
metric_policies: tuple[MetricPolicy, ...],
has_failed: bool,
) -> dict[str, Any]:
"""Format a single benchmark result as a row for the Markdown table."""
metric_values: dict[str, dict[str, float | None]] = {}
metric_values: dict[str, dict[str, Any]] = {}
regressions: list[float] = []
for metric, _label, _precision in METRICS:
curr = safe_float(record.get(metric))
baseline = _baseline_metric(baseline_records, metric)
regression = _metric_delta_percent(metric, record, baseline_records)
metric_values[metric] = {
failing_metrics: list[str] = []
threshold_exceeded_metrics: list[str] = []
for policy in metric_policies:
curr = safe_float(record.get(policy.key))
baseline = _baseline_metric(baseline_records, policy.key)
delta = (
regression_delta(policy, curr, baseline)
if curr is not None and baseline is not None
else None
)
regression = None if delta is None else delta.percent * 100.0
absolute_delta = None if delta is None else delta.absolute
metric_values[policy.key] = {
"curr": curr,
"base": baseline,
"regression_pct": regression,
"absolute_delta": absolute_delta,
"threshold_percent": policy.threshold_percent * 100.0,
"threshold_absolute": policy.threshold_absolute,
"gated": policy.gated,
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
"regressed": False if delta is None else delta.regressed,
}
if regression is not None:
regressions.append(regression)
if delta is not None and delta.threshold_exceeded:
threshold_exceeded_metrics.append(policy.key)
if delta is not None and delta.regressed:
failing_metrics.append(policy.key)
worst_regression_pct = max(regressions) if regressions else None
@@ -313,40 +319,50 @@ def _build_summary_row(
"baseline_n": len(baseline_records),
"metrics": metric_values,
"worst_regression_pct": worst_regression_pct,
"threshold_exceeded_metrics": threshold_exceeded_metrics,
"failing_metrics": failing_metrics,
"failed": has_failed,
}
def _build_markdown_summary(
summary_rows: list[dict[str, Any]],
max_regression: float,
metric_policies: tuple[MetricPolicy, ...],
) -> str:
lines = [
"## Performance Baseline Comparison",
"",
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
"Threshold: gated metrics fail only when both percent and absolute "
"regression floors are exceeded.",
"",
("| Model | GPU | Baseline N | Latency (curr/base) | "
"Throughput (curr/base) | Memory (curr/base) | "
"Text Enc (curr/base) | DiT (curr/base) | "
"VAE Decode (curr/base) | Worst Regression | Status |"),
"|---|---|---:|---|---|---|---|---|---|---:|---|",
"VAE Decode (curr/base) | Worst Regression | Exceeded Metrics | "
"Failing Metrics | Status |"),
"|---|---|---:|---|---|---|---|---|---|---:|---|---|---|",
]
for row in summary_rows:
metric_cells = []
for metric, _label, precision in METRICS:
values = row["metrics"][metric]
metric_cells.append(f"{_compact_value(values['curr'], precision)} / "
f"{_compact_value(values['base'], precision)}")
for policy in metric_policies:
values = row["metrics"][policy.key]
metric_cells.append(f"{_compact_value(values['curr'], policy.precision)} / "
f"{_compact_value(values['base'], policy.precision)}")
worst_reg = ("n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%")
exceeded_metrics = (
", ".join(row["threshold_exceeded_metrics"])
if row["threshold_exceeded_metrics"]
else "none"
)
failing_metrics = ", ".join(row["failing_metrics"]) if row["failing_metrics"] else "none"
status = "FAIL" if row["failed"] else "PASS"
lines.append(f"| {row['model_id']} | {row['gpu_type']} | "
f"{row['baseline_n']} | "
f"{' | '.join(metric_cells)} | "
f"{worst_reg} | {status} |")
f"{worst_reg} | {exceeded_metrics} | {failing_metrics} | {status} |")
return "\n".join(lines) + "\n"
@@ -400,6 +416,7 @@ def main() -> int:
for raw in current_results:
record = _normalize_record(raw)
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
baseline_records = load_records_for_model(
TRACKING_ROOT,
@@ -416,7 +433,7 @@ def main() -> int:
failures: list[str] = []
record["success"] = True
else:
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
failures = _check_regressions(record, baseline_records, metric_policies)
if static_threshold_failed:
failures.append(f"{record['model_id']} fixed-threshold phase failed "
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
@@ -434,10 +451,10 @@ def main() -> int:
print("Tracking upload skipped for "
f"{record['model_id']} ({record['run_source']}, success={record['success']})")
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
summary_rows.append(_build_summary_row(record, baseline_records, metric_policies, bool(failures)))
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
markdown = _build_markdown_summary(summary_rows, resolve_metric_policies(None))
_emit_markdown_summary(markdown, commit_sha)
if all_failures:
+8 -1
View File
@@ -1,12 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
import os
import sys
from html import escape
from datetime import datetime
import plotly.express as px
import pandas as pd
from hf_store import sync_from_hf, load_as_dataframe
try:
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
except ImportError:
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
if repo_root not in sys.path:
sys.path.insert(0, repo_root)
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
TRACKING_ROOT = os.environ.get(
"PERFORMANCE_TRACKING_ROOT",
@@ -0,0 +1,116 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
from fastvideo.tests.performance.test_inference_performance import (
_benchmark_display_id,
_config_identity_metadata,
_is_v2_config,
_validate_benchmark_config,
)
def _v2_config():
return {
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
}
def test_v1_benchmark_config_without_schema_version_validates():
cfg = {
"benchmark_id": "legacy-benchmark",
}
_validate_benchmark_config(cfg, "legacy.json")
assert _is_v2_config(cfg) is False
assert _config_identity_metadata(cfg) == {}
assert _benchmark_display_id(cfg) == "legacy-benchmark"
def test_v2_benchmark_config_identity_validates_and_is_preserved():
cfg = _v2_config()
cfg["quality_metadata"] = {"some": "data"}
_validate_benchmark_config(cfg, "wan.json")
assert _is_v2_config(cfg) is True
assert _config_identity_metadata(cfg) == {
"config_schema_version": 2,
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
"quality_metadata": {"some": "data"},
}
def test_v2_benchmark_config_missing_identity_fields_fails_clearly():
cfg = _v2_config()
del cfg["variant_id"]
del cfg["benchmark_version"]
expected = "wan.json: missing required v2 identity fields: variant_id, benchmark_version"
with pytest.raises(ValueError, match=expected):
_validate_benchmark_config(cfg, "wan.json")
@pytest.mark.parametrize(
("field", "value"),
[
("workload_id", {}),
("workload_id", ""),
("workload_id", " "),
("variant_id", []),
("variant_id", ""),
("variant_id", " "),
],
)
def test_v2_benchmark_config_rejects_invalid_string_identity_values(field, value):
cfg = _v2_config()
cfg[field] = value
expected = f"wan.json: v2 identity field {field!r} must be a non-empty string"
with pytest.raises(ValueError, match=expected):
_validate_benchmark_config(cfg, "wan.json")
@pytest.mark.parametrize("value", [None, "1", 1.5, True])
def test_v2_benchmark_config_rejects_invalid_benchmark_version_values(value):
cfg = _v2_config()
cfg["benchmark_version"] = value
expected = "wan.json: v2 identity field 'benchmark_version' must be an integer"
with pytest.raises(ValueError, match=expected):
_validate_benchmark_config(cfg, "wan.json")
def test_partial_v2_identity_requires_schema_version():
cfg = {
"benchmark_id": "wan-t2v-1.3b-2gpu",
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
}
expected = "wan.json: v2 benchmark identity fields require config_schema_version=2"
with pytest.raises(ValueError, match=expected):
_validate_benchmark_config(cfg, "wan.json")
def test_optional_v2_metadata_fields_must_be_objects():
cfg = {
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v-1.3b",
"variant_id": "canonical",
"benchmark_version": 1,
"quality_metadata": ["not", "an", "object"],
}
expected = "wan.json: optional v2 metadata field 'quality_metadata' must be an object"
with pytest.raises(ValueError, match=expected):
_validate_benchmark_config(cfg, "wan.json")
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.tests.performance import compare_baseline
from fastvideo.performance.metric_policy import resolve_metric_policies
def _raw_result():
@@ -72,9 +73,157 @@ def test_normalized_record_includes_source_metadata(monkeypatch):
assert record["job_id"] == "job-1"
def test_normalized_record_includes_effective_regression_thresholds(monkeypatch):
monkeypatch.setenv("PERF_RUN_SOURCE", "pr")
raw = _raw_result()
raw["regression_thresholds"] = {
"latency": {
"threshold_percent": 0.09,
"threshold_absolute": 0.75,
"gated": True,
},
"throughput": {
"gated": False,
},
}
record = compare_baseline.normalize_performance_result(raw)
assert record["regression_thresholds"]["latency"] == {
"threshold_percent": 0.09,
"threshold_absolute": 0.75,
"gated": True,
}
assert record["regression_thresholds"]["throughput"]["gated"] is False
def test_invalid_regression_threshold_container_uses_defaults():
policies = resolve_metric_policies(["not", "a", "mapping"])
latency = next(policy for policy in policies if policy.key == "latency")
assert latency.threshold_percent == 0.08
assert latency.threshold_absolute == 0.5
assert latency.gated is True
def test_boolean_regression_threshold_values_are_ignored():
policies = resolve_metric_policies({
"latency": {
"threshold_percent": True,
"threshold_absolute": False,
"gated": "false",
}
})
latency = next(policy for policy in policies if policy.key == "latency")
assert latency.threshold_percent == 0.08
assert latency.threshold_absolute == 0.5
assert latency.gated is False
def test_baseline_eligibility_only_for_successful_scheduled_main():
assert compare_baseline._is_baseline_eligible("scheduled_main", True) is True
assert compare_baseline._is_baseline_eligible("scheduled_main", False) is False
assert compare_baseline._is_baseline_eligible("pr", True) is False
assert compare_baseline._is_baseline_eligible("local", True) is False
def test_latency_regression_requires_percent_and_absolute_floors():
baseline = [{"latency": 10.0}]
current = {"model_id": "wan", "latency": 10.6}
percent_only = resolve_metric_policies({
"latency": {
"threshold_percent": 0.05,
"threshold_absolute": 0.75,
}
})
absolute_only = resolve_metric_policies({
"latency": {
"threshold_percent": 0.10,
"threshold_absolute": 0.5,
}
})
both = resolve_metric_policies({
"latency": {
"threshold_percent": 0.05,
"threshold_absolute": 0.5,
}
})
assert compare_baseline._check_regressions(current, baseline, percent_only) == []
assert compare_baseline._check_regressions(current, baseline, absolute_only) == []
failures = compare_baseline._check_regressions(current, baseline, both)
assert len(failures) == 1
assert "latency regressed by 6.0% and 0.600" in failures[0]
def test_throughput_regression_uses_higher_is_better_direction():
baseline = [{"throughput": 10.0}]
current = {"model_id": "wan", "throughput": 9.0}
policies = resolve_metric_policies({
"throughput": {
"threshold_percent": 0.05,
"threshold_absolute": 0.5,
}
})
failures = compare_baseline._check_regressions(current, baseline, policies)
assert len(failures) == 1
assert "throughput regressed by 10.0% and 1.000" in failures[0]
def test_memory_regression_uses_metric_specific_absolute_floor():
baseline = [{"memory": 10000.0}]
current = {"model_id": "wan", "memory": 10600.0}
policies = resolve_metric_policies({
"memory": {
"threshold_percent": 0.05,
"threshold_absolute": 256.0,
}
})
failures = compare_baseline._check_regressions(current, baseline, policies)
assert len(failures) == 1
assert "memory regressed by 6.0% and 600.0" in failures[0]
def test_component_metric_can_gate_independently():
baseline = [{"dit_time_s": 8.0}]
current = {"model_id": "wan", "dit_time_s": 8.6}
policies = resolve_metric_policies({
"dit_time_s": {
"threshold_percent": 0.05,
"threshold_absolute": 0.25,
}
})
failures = compare_baseline._check_regressions(current, baseline, policies)
assert len(failures) == 1
assert "dit_time_s regressed by 7.5% and 0.600" in failures[0]
def test_informational_metric_remains_visible_without_failing():
baseline = [{"throughput": 10.0}]
current = {"model_id": "wan", "gpu_type": "NVIDIA L40S", "throughput": 8.0}
policies = resolve_metric_policies({
"throughput": {
"threshold_percent": 0.01,
"threshold_absolute": 0.01,
"gated": False,
}
})
row = compare_baseline._build_summary_row(current, baseline, policies, False)
assert compare_baseline._check_regressions(current, baseline, policies) == []
assert row["metrics"]["throughput"]["regression_pct"] == 20.0
assert row["metrics"]["throughput"]["gated"] is False
assert row["metrics"]["throughput"]["threshold_exceeded"] is True
assert row["metrics"]["throughput"]["regressed"] is False
assert row["threshold_exceeded_metrics"] == ["throughput"]
assert row["failing_metrics"] == []
@@ -68,6 +68,8 @@ def test_summary_endpoint_returns_latest_group_status():
assert body["count"] == 1
assert body["status_counts"] == {"pass": 1, "fail": 0}
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
assert body["rows"][0]["metrics"]["latency"]["threshold_exceeded"] is True
assert body["rows"][0]["threshold_exceeded_metrics"] == ["latency", "throughput"]
assert body["rows"][0]["computed_regression_status"] == "fail"
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.performance import hf_store
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
from fastvideo.tests.performance import hf_store
def _record(ts, commit, latency, throughput, success=True, **metadata):
@@ -28,16 +28,23 @@ def test_build_latest_summary_uses_previous_successful_records_for_baseline():
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
]
rows = build_latest_summary(records, max_regression=0.05)
rows = build_latest_summary(records)
assert len(rows) == 1
row = rows[0]
assert row["baseline_n"] == 1
assert row["metrics"]["latency"]["baseline"] == 10.0
assert row["metrics"]["latency"]["regression_pct"] == 10.0
assert row["metrics"]["latency"]["absolute_delta"] == 1.0
assert row["metrics"]["latency"]["threshold_percent"] == 8.0
assert row["metrics"]["latency"]["threshold_absolute"] == 0.5
assert row["metrics"]["latency"]["threshold_exceeded"] is True
assert row["metrics"]["latency"]["regressed"] is True
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
assert row["status"] == "pass"
assert row["computed_regression_status"] == "fail"
assert row["threshold_exceeded_metrics"] == ["latency", "throughput"]
assert row["failing_metrics"] == ["latency", "throughput"]
def test_build_latest_summary_status_uses_latest_record_success_field():
@@ -73,7 +80,7 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
),
]
rows = build_latest_summary(records, max_regression=0.05, run_source="pr")
rows = build_latest_summary(records, run_source="pr")
assert len(rows) == 1
assert rows[0]["run_source"] == "pr"
@@ -83,6 +90,60 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
assert rows[0]["computed_regression_status"] == "fail"
def test_build_latest_summary_requires_absolute_floor_for_computed_regression():
records = [
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
_record(
"2026-01-02T00:00:00+00:00",
"b" * 40,
10.6,
10.0,
regression_thresholds={
"latency": {
"threshold_percent": 0.05,
"threshold_absolute": 0.75,
"gated": True,
}
},
),
]
rows = build_latest_summary(records)
assert round(rows[0]["metrics"]["latency"]["regression_pct"], 1) == 6.0
assert round(rows[0]["metrics"]["latency"]["absolute_delta"], 3) == 0.6
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is False
assert rows[0]["metrics"]["latency"]["regressed"] is False
assert rows[0]["computed_regression_status"] == "pass"
def test_build_latest_summary_separates_informational_threshold_crossing():
records = [
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
_record(
"2026-01-02T00:00:00+00:00",
"b" * 40,
10.6,
10.0,
regression_thresholds={
"latency": {
"threshold_percent": 0.05,
"threshold_absolute": 0.5,
"gated": False,
}
},
),
]
rows = build_latest_summary(records)
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is True
assert rows[0]["metrics"]["latency"]["regressed"] is False
assert rows[0]["threshold_exceeded_metrics"] == ["latency"]
assert rows[0]["failing_metrics"] == []
assert rows[0]["computed_regression_status"] == "pass"
def test_filter_records_and_trends_preserve_metric_points():
records = [
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
@@ -28,6 +28,17 @@ STAGE_METRIC_MAP: dict[str, str] = {
"DmdDenoisingStage": "dit_time_s",
"DecodingStage": "vae_decode_time_s",
}
V2_CONFIG_SCHEMA_VERSION = 2
V2_REQUIRED_IDENTITY_FIELDS = (
"workload_id",
"variant_id",
"benchmark_version",
)
V2_OPTIONAL_METADATA_FIELDS = (
"recipe",
"metric_threshold_policy",
"quality_metadata",
)
# -- Config discovery -------------------------------------------------------
@@ -42,6 +53,71 @@ _BENCHMARKS_DIR = os.path.join(
)
def _has_v2_fields(cfg):
v2_fields = V2_REQUIRED_IDENTITY_FIELDS + V2_OPTIONAL_METADATA_FIELDS
return any(field in cfg for field in v2_fields)
def _is_v2_config(cfg):
return cfg.get("config_schema_version") == V2_CONFIG_SCHEMA_VERSION
def _validate_non_empty_string(value, field, path):
if not isinstance(value, str) or not value.strip():
raise ValueError(f"{path}: v2 identity field {field!r} must be a non-empty string")
def _validate_integer(value, field, path):
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"{path}: v2 identity field {field!r} must be an integer")
def _validate_benchmark_config(cfg, path="<memory>"):
missing_common = [field for field in ("benchmark_id",) if field not in cfg]
if missing_common:
raise ValueError(f"{path}: missing required benchmark config fields: {', '.join(missing_common)}")
schema_version = cfg.get("config_schema_version")
if schema_version is None:
if _has_v2_fields(cfg):
raise ValueError(f"{path}: v2 benchmark identity fields require config_schema_version=2")
return
if schema_version != V2_CONFIG_SCHEMA_VERSION:
raise ValueError(f"{path}: unsupported benchmark config_schema_version={schema_version!r}")
missing_v2 = [field for field in V2_REQUIRED_IDENTITY_FIELDS if field not in cfg]
if missing_v2:
raise ValueError(f"{path}: missing required v2 identity fields: {', '.join(missing_v2)}")
_validate_non_empty_string(cfg["workload_id"], "workload_id", path)
_validate_non_empty_string(cfg["variant_id"], "variant_id", path)
_validate_integer(cfg["benchmark_version"], "benchmark_version", path)
for field in V2_OPTIONAL_METADATA_FIELDS:
if field in cfg and not isinstance(cfg[field], Mapping):
raise ValueError(f"{path}: optional v2 metadata field {field!r} must be an object")
def _config_identity_metadata(cfg):
if not _is_v2_config(cfg):
return {}
metadata = {
"config_schema_version": cfg["config_schema_version"],
"workload_id": cfg["workload_id"],
"variant_id": cfg["variant_id"],
"benchmark_version": cfg["benchmark_version"],
}
for field in V2_OPTIONAL_METADATA_FIELDS:
if field in cfg:
metadata[field] = cfg[field]
return metadata
def _benchmark_display_id(cfg):
return cfg["benchmark_id"]
def _discover_benchmarks():
"""Glob benchmark JSON configs and return list of (id, config) tuples."""
pattern = os.path.join(_BENCHMARKS_DIR, "*.json")
@@ -49,6 +125,7 @@ def _discover_benchmarks():
for path in sorted(glob.glob(pattern)):
with open(path) as f:
cfg = json.load(f)
_validate_benchmark_config(cfg, path)
configs.append(cfg)
return configs
@@ -102,7 +179,11 @@ def _extract_component_times(result: dict) -> dict[str, float | None]:
logger.debug("Skipping malformed stage '%s' data: %r", stage_name, stage_data)
continue
stage_class = stage_data.get("stage_class", stage_name)
metric_key = STAGE_METRIC_MAP.get(stage_class)
component_metric = stage_data.get("component_metric")
if isinstance(component_metric, str) and component_metric in component_times:
metric_key = component_metric
else:
metric_key = STAGE_METRIC_MAP.get(stage_class)
if metric_key is None:
logger.debug("Unmapped stage '%s' class '%s' (%.3fs)",
stage_name,
@@ -219,6 +300,7 @@ def _run_benchmark(cfg):
results = {
"benchmark_id": cfg["benchmark_id"],
**_config_identity_metadata(cfg),
"model_short_name": model_info.get("model_short_name", ""),
"device": device_name,
"num_gpus": init_kwargs.get("num_gpus", 1),
@@ -231,6 +313,7 @@ def _run_benchmark(cfg):
"max_peak_memory_mb": round(max_peak_memory, 1),
"individual_peak_memories_mb": [round(m, 1) for m in peak_memories],
"thresholds": thresholds,
"regression_thresholds": cfg.get("regression_thresholds", {}),
"commit": os.environ.get("BUILDKITE_COMMIT", ""),
"pr_number": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"timestamp": datetime.now(timezone.utc).isoformat(),
@@ -275,7 +358,7 @@ def _run_benchmark(cfg):
@pytest.mark.parametrize(
"cfg",
_BENCHMARK_CONFIGS,
ids=[c["benchmark_id"] for c in _BENCHMARK_CONFIGS],
ids=[_benchmark_display_id(c) for c in _BENCHMARK_CONFIGS],
)
def test_inference_performance(cfg):
"""Measure generation latency, peak GPU memory, and component-level timings
@@ -1,13 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.pipelines.pipeline_batch_info import PipelineLoggingInfo
from fastvideo.pipelines.stages.denoising import Cosmos25AutoDenoisingStage, DenoisingStage
from fastvideo.pipelines.stages.text_encoding import Cosmos25TextEncodingStage
from fastvideo.tests.performance.test_inference_performance import _extract_component_times
class SubclassStyleDenoisingStage(DenoisingStage):
pass
def test_extract_component_times_handles_pipeline_logging_info_object():
logging_info = PipelineLoggingInfo()
logging_info.add_stage_execution_time("prompt_encoding_stage", 1.25)
logging_info.add_stage_metric("prompt_encoding_stage", "stage_class", "TextEncodingStage")
logging_info.add_stage_metric("prompt_encoding_stage", "component_metric", "text_encoder_time_s")
assert _extract_component_times({"logging_info": logging_info}) == {
"text_encoder_time_s": 1.25,
@@ -45,6 +52,35 @@ def test_extract_component_times_uses_stage_class_for_pipeline_stage_keys():
}
def test_denoising_stage_subclasses_inherit_component_metric():
assert SubclassStyleDenoisingStage.performance_component_metric == "dit_time_s"
def test_cosmos25_direct_pipeline_stages_define_component_metrics():
assert Cosmos25TextEncodingStage.performance_component_metric == "text_encoder_time_s"
assert Cosmos25AutoDenoisingStage.performance_component_metric == "dit_time_s"
def test_extract_component_times_uses_component_metric_for_stage_subclasses():
result = {
"logging_info": {
"stages": {
"denoising_stage": {
"execution_time": 4.2,
"stage_class": "CosmosDenoisingStage",
"component_metric": "dit_time_s",
},
},
},
}
assert _extract_component_times(result) == {
"text_encoder_time_s": None,
"dit_time_s": 4.2,
"vae_decode_time_s": None,
}
def test_extract_component_times_keeps_legacy_class_name_keys():
# Backward-compatibility check for logs produced before pipeline-unique
# stage keys carried a separate stage_class field.
@@ -3,9 +3,9 @@
from __future__ import annotations
import os
from collections.abc import Iterator
from contextlib import contextmanager
from logging import Logger
from typing import Iterator
from fastvideo import VideoGenerator
from fastvideo.tests.ssim.reference_utils import (
@@ -67,6 +67,12 @@ def _find_reference_video(reference_folder: str, prompt: str) -> str:
raise FileNotFoundError("Reference video missing")
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
stale_path = os.path.join(output_dir, output_video_name)
if os.path.exists(stale_path):
os.remove(stale_path)
def _assert_similarity(
*,
logger: Logger,
@@ -214,6 +220,7 @@ def run_text_to_video_similarity_test(
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
params_map = select_ssim_params(
default_params_map,
@@ -289,6 +296,7 @@ def run_image_to_video_similarity_test(
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
params_map = select_ssim_params(
default_params_map,
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
from pathlib import Path
import pytest
from PIL import Image, ImageDraw
from fastvideo.logger import init_logger
from fastvideo.tests.ssim.inference_similarity_utils import (
resolve_inference_device_reference_folder,
run_image_to_video_similarity_test,
run_text_to_video_similarity_test,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 1
device_reference_folder = resolve_inference_device_reference_folder(logger)
_LOCAL_CONVERTED_MODEL = Path("converted_weights/dreamx_world")
_MODEL_PATH = os.getenv(
"DREAMX_WORLD_SSIM_MODEL_PATH",
str(_LOCAL_CONVERTED_MODEL),
)
_LOCAL_AR_CANDIDATES = (
Path("/tmp/converted_dreamx_world_ar"),
Path("/root/data/dreamx_world_ar_converted"),
)
_DEFAULT_AR_MODEL_PATH = next(
(str(path) for path in _LOCAL_AR_CANDIDATES if path.exists()),
str(_LOCAL_AR_CANDIDATES[0]),
)
_AR_MODEL_PATH = os.getenv(
"DREAMX_WORLD_AR_SSIM_MODEL_PATH",
_DEFAULT_AR_MODEL_PATH,
)
DREAMX_WORLD_PARAMS = {
"num_gpus": 1,
"model_path": _MODEL_PATH,
"height": 64,
"width": 64,
"num_frames": 9,
"num_inference_steps": 1,
"guidance_scale": 1.0,
"seed": 1024,
"fps": 16,
}
DREAMX_WORLD_FULL_QUALITY_PARAMS = {
**DREAMX_WORLD_PARAMS,
"height": 480,
"width": 832,
"num_frames": 161,
"num_inference_steps": 30,
"guidance_scale": 5.0,
}
DREAMX_WORLD_AR_PARAMS = {
"num_gpus": 1,
"model_path": _AR_MODEL_PATH,
"height": 192,
"width": 192,
"num_frames": 81,
"num_inference_steps": 4,
"guidance_scale": 1.0,
"seed": 2048,
"fps": 16,
}
DREAMX_WORLD_AR_FULL_QUALITY_PARAMS = {
**DREAMX_WORLD_AR_PARAMS,
"height": 704,
"width": 1280,
"num_frames": 1005,
}
DREAMX_WORLD_MODEL_TO_PARAMS = {
"DreamX-World-5B-Cam": DREAMX_WORLD_PARAMS,
}
DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
"DreamX-World-5B": DREAMX_WORLD_AR_PARAMS,
}
FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS = {
"DreamX-World-5B-Cam": DREAMX_WORLD_FULL_QUALITY_PARAMS,
}
FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
"DreamX-World-5B": DREAMX_WORLD_AR_FULL_QUALITY_PARAMS,
}
DREAMX_WORLD_TEST_CASES = [
(
"A cinematic first-person drive through a futuristic coastal city at sunrise, "
"reflective glass towers, clean streets, soft volumetric light.",
("w", "d", "w"),
(4.0, 2.0, 4.0),
),
]
DREAMX_WORLD_AR_TEST_CASES = [
(
"A long autonomous drive through a futuristic coastal city at sunrise, "
"smooth forward camera motion, reflective glass towers, clean streets.",
("w", "d", "w", "a"),
(2.0, 1.0, 2.0, 1.0),
),
]
def _write_deterministic_reference_image(path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
draw = ImageDraw.Draw(image)
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
draw.line((0, 80, 96, 66), fill=(238, 209, 124), width=3)
image.save(path)
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_TEST_CASES)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_MODEL_TO_PARAMS.keys()))
def test_dreamx_world_inference_similarity(
prompt: str,
action_list: tuple[str, ...],
action_speed_list: tuple[float, ...],
attention_backend_name: str,
model_id: str,
tmp_path: Path,
) -> None:
model_path = Path(str(DREAMX_WORLD_MODEL_TO_PARAMS[model_id]["model_path"]))
if not model_path.exists():
pytest.skip(
f"DreamX-World converted model path is missing: {model_path}. "
"Set DREAMX_WORLD_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
)
image_path = tmp_path / "dreamx_world_ssim_input.png"
_write_deterministic_reference_image(image_path)
run_image_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
image_path=str(image_path),
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=DREAMX_WORLD_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS,
min_acceptable_ssim=0.98,
init_kwargs_override={
"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",
},
generation_kwargs_override={
"action_list": list(action_list),
"action_speed_list": list(action_speed_list),
},
)
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_AR_TEST_CASES)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_AR_MODEL_TO_PARAMS.keys()))
def test_dreamx_world_ar_inference_similarity(
prompt: str,
action_list: tuple[str, ...],
action_speed_list: tuple[float, ...],
attention_backend_name: str,
model_id: str,
) -> None:
model_path = Path(str(DREAMX_WORLD_AR_MODEL_TO_PARAMS[model_id]["model_path"]))
if not model_path.exists():
pytest.skip(
f"DreamX-World AR converted model path is missing: {model_path}. "
"Set DREAMX_WORLD_AR_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
)
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=DREAMX_WORLD_AR_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS,
min_acceptable_ssim=0.98,
init_kwargs_override={
"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": "DreamXWorldARPipeline",
},
generation_kwargs_override={
"action_list": list(action_list),
"action_speed_list": list(action_speed_list),
},
)
@@ -21,6 +21,26 @@ class FakeTokenizer:
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
})
class FakeChatTokenizer:
def __init__(self):
self.last_messages = None
self.last_kwargs = None
def apply_chat_template(self, messages, **kwargs):
self.last_messages = messages
self.last_kwargs = kwargs
assert isinstance(messages[0], list)
assert messages[0][0]["role"] == "system"
assert messages[0][1]["role"] == "user"
B = len(messages)
seq_len = int(kwargs.get("max_length", 4))
return TensorDict({
"input_ids": torch.arange(B * seq_len).view(B, seq_len),
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
})
class FakeTextEncoder(torch.nn.Module):
def __init__(self, hidden_size=8):
super().__init__()
@@ -38,6 +58,14 @@ class FakeTextEncoder(torch.nn.Module):
def id_preprocess(x: str) -> str:
return x
def chat_list_preprocess(x: str):
return [
{"role": "system", "content": "Describe the video."},
{"role": "user", "content": x if x else " "},
]
def take_mean_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
# [B, T, H] -> [B, H]
return outputs.last_hidden_state.mean(dim=1)
@@ -156,3 +184,32 @@ def test_encode_text_does_not_force_hidden_states_for_ltx2_prefix():
stage.encode_text("a", fastvideo_args, encoder_index=[0])
assert stage.text_encoders[0].last_output_hidden_states is False
def test_chat_list_preprocess_output_is_not_stripped():
fastvideo_args, hidden = make_args(num_encoders=1, text_len=5, hidden_size=8)
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
encoder_config.is_chat_model = True
encoder_config.treat_empty_as_dot = True
fastvideo_args.pipeline_config.preprocess_text_funcs = (chat_list_preprocess, )
tokenizer = FakeChatTokenizer()
stage = TextEncodingStage(
text_encoders=[FakeTextEncoder(hidden_size=hidden)],
tokenizers=[tokenizer],
)
embeds, masks = stage.encode_text(
"a robotic arm welding a metal structure",
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
assert embeds[0].shape == (1, hidden)
assert masks[0].shape == (1, 5)
assert tokenizer.last_messages == [[
{"role": "system", "content": "Describe the video."},
{"role": "user", "content": "a robotic arm welding a metal structure"},
]]
assert tokenizer.last_kwargs["return_tensors"] == "pt"
@@ -127,4 +127,12 @@ def test_wan_causal_dfsft_single_train_step(
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
# rtol above the harness default: the causal model compiles flex_attention
# with max-autotune (required for Wan 1.3B's head config), and the
# timing-based kernel selection is bimodal across L40S containers —
# observed 3.2562 vs 3.5860 (10.13% apart) with identical code, straddling
# the default 10%. 12% covers both winners; real wiring breakage (dead
# grads, scale bugs) still lands far outside it.
check_grad_norm_regression("test_wan_causal_dfsft",
model.transformer,
rtol=0.12)
@@ -0,0 +1,124 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Convert DreamX-World-5B autoregressive weights to FastVideo layout.
The HF repository stores one raw official ``model.safetensors`` whose keys match
FastVideo's native ``DreamXWorldARTransformer3DModel``. The converter writes a
Diffusers-like root with ``transformer/config.json`` and reusable Wan2.2
components. Use ``--symlink-transformer`` locally to avoid duplicating the 21GB
AR tensor file.
"""
from __future__ import annotations
import argparse
import json
import shutil
from pathlib import Path
TRANSFORMER_CONFIG: dict[str, object] = {
"_class_name": "DreamXWorldARTransformer3DModel",
"model_type": "ti2v",
"patch_size": [1, 2, 2],
"text_len": 512,
"num_attention_heads": 24,
"attention_head_dim": 128,
"in_channels": 48,
"out_channels": 48,
"text_dim": 4096,
"freq_dim": 256,
"ffn_dim": 14336,
"num_layers": 30,
"local_attn_size": 12,
"sink_size": 3,
"cross_attn_norm": True,
"qk_norm": True,
"eps": 1e-6,
"add_control_adapter": True,
"cam_method": "prope",
"attn_compress": 4,
"cam_self_attn_layers": list(range(30)),
"num_frames_per_block": 3,
}
MODEL_INDEX: dict[str, object] = {
"_class_name": "DreamXWorldARPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldARTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
def _source_safetensors(source: Path) -> Path:
if source.is_file():
return source
path = source / "model.safetensors"
if not path.exists():
raise FileNotFoundError(f"Missing AR model.safetensors under {source}")
return path
def convert_transformer(source: Path, output: Path, symlink_transformer: bool) -> None:
src = _source_safetensors(source)
transformer_dir = output / "transformer"
transformer_dir.mkdir(parents=True, exist_ok=True)
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
dst = transformer_dir / "model.safetensors"
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
return
if symlink_transformer:
dst.symlink_to(src.resolve())
else:
shutil.copy2(src, dst)
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
src = component_source / component
dst = output / component
if not src.exists():
raise FileNotFoundError(f"Missing reused component source: {src}")
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
return
if symlink:
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
elif src.is_dir():
shutil.copytree(src, dst)
else:
shutil.copy2(src, dst)
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
if component_source is not None:
for component in REUSED_COMPONENTS:
_copy_or_link_component(component, component_source, output, symlink_components)
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
if missing:
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
return
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--component-source", type=Path)
parser.add_argument("--symlink-components", action="store_true")
parser.add_argument("--symlink-transformer", action="store_true")
args = parser.parse_args()
convert_transformer(args.source, args.output, args.symlink_transformer)
write_model_index(args.output, args.component_source, args.symlink_components)
if __name__ == "__main__":
main()
@@ -0,0 +1,221 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Convert DreamX-World-5B-Cam raw transformer weights to FastVideo-loadable format.
The GD-ML/DreamX-World-5B-Cam repository stores the transformer as raw
DreamX/Wan official shards. FastVideo's TransformerLoader expects a Diffusers-like
transformer folder with a config.json and safetensors whose keys can be mapped by
WanVideoConfig.param_names_mapping. This script performs the raw official ->
Diffusers-like key rename and writes the DreamX 5B-Cam transformer config.
Example:
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
--source official_weights/dreamx_world \
--output converted_weights/dreamx_world
"""
from __future__ import annotations
import argparse
import json
import re
import shutil
from collections import OrderedDict
from pathlib import Path
import torch
from huggingface_hub import save_torch_state_dict
from safetensors import safe_open
from safetensors.torch import load_file
OFFICIAL_TO_DIFFUSERS_MAPPING: dict[str, str] = {
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.linear_1.\1",
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.linear_2.\1",
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.linear_1.\1",
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.linear_2.\1",
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_proj.\1",
r"^img_emb\.proj\.0\.(.*)$": r"condition_embedder.image_embedder.norm1.\1",
r"^img_emb\.proj\.1\.(.*)$": r"condition_embedder.image_embedder.ff.net.0.proj.\1",
r"^img_emb\.proj\.3\.(.*)$": r"condition_embedder.image_embedder.ff.net.2.\1",
r"^img_emb\.proj\.4\.(.*)$": r"condition_embedder.image_embedder.norm2.\1",
r"^head\.modulation": r"scale_shift_table",
r"^head\.head\.(.*)$": r"proj_out.\1",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.attn1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.attn1.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.k_img\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.v_img\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.net.0.proj.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.net.2.\2",
r"^blocks\.(\d+)\.modulation": r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm2.\2",
}
TRANSFORMER_CONFIG: dict[str, object] = {
"_class_name": "DreamXWorldTransformer3DModel",
"patch_size": [1, 2, 2],
"text_len": 512,
"num_attention_heads": 24,
"attention_head_dim": 128,
"in_channels": 48,
"out_channels": 48,
"text_dim": 4096,
"freq_dim": 256,
"ffn_dim": 14336,
"num_layers": 30,
"cross_attn_norm": True,
"qk_norm": "rms_norm_across_heads",
"eps": 1e-6,
"image_dim": None,
"added_kv_proj_dim": None,
"rope_max_seq_len": 1024,
"add_control_adapter": True,
"cam_method": "prope",
"attn_compress": 1,
"cam_self_attn_layers": None,
}
MODEL_INDEX: dict[str, object] = {
"_class_name": "DreamXWorldPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
def map_transformer_key(key: str) -> str:
for pattern, replacement in OFFICIAL_TO_DIFFUSERS_MAPPING.items():
if re.match(pattern, key):
return re.sub(pattern, replacement, key)
return key
def _safetensor_files(source: Path) -> list[Path]:
if source.is_file():
if source.suffix != ".safetensors":
raise ValueError(f"Only .safetensors files are supported, got {source}")
return [source]
index_path = source / "diffusion_pytorch_model.safetensors.index.json"
if index_path.exists():
index = json.loads(index_path.read_text())
return sorted({source / shard for shard in index["weight_map"].values()})
files = sorted(source.glob("*.safetensors"))
if not files:
raise FileNotFoundError(f"No safetensors files found under {source}")
return files
def convert_transformer(source: Path, output: Path, max_shard_size: str) -> None:
transformer_dir = output / "transformer"
transformer_dir.mkdir(parents=True, exist_ok=True)
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
for shard in _safetensor_files(source):
print(f"loading {shard}")
for key, tensor in load_file(shard, device="cpu").items():
new_key = map_transformer_key(key)
if new_key in converted:
raise ValueError(f"Duplicate converted key: {new_key}")
converted[new_key] = tensor
print(f"saving {len(converted)} tensors to {transformer_dir}")
save_torch_state_dict(converted, transformer_dir, max_shard_size=max_shard_size)
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
src = component_source / component
dst = output / component
if not src.exists():
raise FileNotFoundError(f"Missing reused component source: {src}")
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
print(f"keeping existing {dst}")
return
if symlink:
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
print(f"linked {dst} -> {src}")
elif src.is_dir():
shutil.copytree(src, dst)
print(f"copied {src} -> {dst}")
else:
shutil.copy2(src, dst)
print(f"copied {src} -> {dst}")
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
if component_source is not None:
for component in REUSED_COMPONENTS:
_copy_or_link_component(component, component_source, output, symlink_components)
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
if missing:
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
print("pass --component-source <Wan2.2 Diffusers root> to copy or link reused components")
return
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
print(f"wrote {output / 'model_index.json'}")
def analyze(source: Path) -> None:
total = 0
unchanged = 0
examples: list[tuple[str, str]] = []
for shard in _safetensor_files(source):
with safe_open(shard, framework="pt", device="cpu") as tensors:
for key in tensors:
total += 1
new_key = map_transformer_key(key)
unchanged += int(new_key == key)
if len(examples) < 20 and new_key != key:
examples.append((key, new_key))
print(f"total_keys={total} unchanged_keys={unchanged}")
for old, new in examples:
print(f"{old} -> {new}")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", type=Path, required=True, help="DreamX raw transformer directory or safetensors file")
parser.add_argument("--output", type=Path, required=True, help="Output model root; transformer/ is created inside it")
parser.add_argument("--max-shard-size", default="10GB")
parser.add_argument(
"--component-source",
type=Path,
help="Optional Wan2.2 Diffusers root whose scheduler/text_encoder/tokenizer/vae components are reused.",
)
parser.add_argument(
"--symlink-components",
action="store_true",
help="Symlink reused components from --component-source instead of copying them.",
)
parser.add_argument("--analyze", action="store_true", help="Only print key mapping summary")
args = parser.parse_args()
if args.analyze:
analyze(args.source)
else:
convert_transformer(args.source, args.output, args.max_shard_size)
write_model_index(args.output, args.component_source, args.symlink_components)
if __name__ == "__main__":
main()
@@ -0,0 +1,189 @@
# DreamX World Port Status
## Summary
- model_family: `dreamx_world`
- workload_types: `I2V camera-control compatibility shim`; `I2V autoregressive camera-control forcing`
- official_ref: `https://github.com/AMAP-ML/DreamX-World`
- official_ref_dir: `DreamX-World/`
- hf_weights_path: `GD-ML/DreamX-World-5B-Cam`
- local_weights_dir: `official_weights/dreamx_world`
- source_layout: `raw_official`
- local_tests_readme: `tests/local_tests/dreamx_world/README.md`
## Current Phase
- phase: `phase_11_post_parity_handoff`
- status: `complete`
- owner: `orchestrator`
- last_updated: `2026-07-02`
## Component Matrix
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|---|---|---|---|---|---|---|---|---|---|
| transformer | dit | ported_dedicated | `DreamX-World/models/wan_transformer3d.py`; PRoPE helpers in `DreamX-World/models/prope_utils.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `Wan2_2Transformer3DModel.from_pretrained(... cam_method=prope, add_control_adapter=True)` | `fastvideo/models/dits/dreamx_world.py`; `fastvideo/configs/models/dits/dreamx_world.py`; DreamX pipeline config helper | native_prope_pass | real_conversion_pass | strict_load_and_forward_parity_pass | none |
| vae | vae | reuse_pending | `DreamX-World/models/wan_vae3_8.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `AutoencoderKLWan3_8.from_pretrained(Wan2.2_VAE.pth)` | `fastvideo/models/vaes/wanvae.py`; DreamX VAE config helper | config_smoke_pass | raw_key_mapping_pass | encode_parity_pass | none |
| text_encoder/tokenizer | encoder | reuse_pending | `DreamX-World/models/wan_text_encoder.py`; tokenizer via Wan2.2 base model | `DreamX-World/inference_dreamx5b.py::setup_models`, `WanT5EncoderModel` + tokenizer subpaths | `fastvideo/models/encoders/t5.py::UMT5EncoderModel`; DreamX UMT5 config helper | config_smoke_pass | staged_weight_load_pass | hidden_state_parity_pass | none |
| scheduler | generic | reuse_proven | Diffusers `FlowMatchEulerDiscreteScheduler` | `DreamX-World/inference_dreamx5b.py::setup_models`, default `sampler_name=Flow` | `fastvideo/models/schedulers/scheduling_flow_match_euler_discrete.py` | pass | not_required | non_skip_pass | Q003 |
| camera_conditioning | generic | port_pending | `DreamX-World/utils/inference_utils.py`, `DreamX-World/models/prope_utils.py`, `DreamX-World/wan/modules/camera_prope.py` | `DreamX-World/inference_dreamx5b.py::get_camera_sequence`, `pipeline(... control_camera_video=...)` | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` | pass | not_required | non_skip_pass | none |
| pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_dreamxworld.py` | `DreamX-World/inference_dreamx5b.py::process_inference_from_json` | `fastvideo/pipelines/basic/dreamx_world/` plus config/preset/registry | pipeline_load_generate_smoke_pass | model_index_and_config_consistency_smoke_pass | pipeline_api_vs_worker_forward_parity_pass | none |
| ar_transformer | dit | port_complete | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py::CausalWanModel` | `DreamX-World/inference_ar_forcing.py::load_pipeline` | `fastvideo/models/dits/dreamx_world_ar.py`; `fastvideo/configs/models/dits/dreamx_world.py::DreamXWorldARConfig` | tiny_official_forward_parity_pass | identity_conversion_pass | real_5b_strict_load_pass | none |
| ar_pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_causal_camera.py` | `DreamX-World/inference_ar_forcing.py::main` | `fastvideo/pipelines/basic/dreamx_world/dreamx_world_ar_pipeline.py`; `fastvideo/pipelines/basic/dreamx_world/ar_denoising.py`; registry/preset/config | config_registry_pass | symlink_model_index_pass | short_full_generation_pass | none |
## Conversion State
- conversion_script: `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py`
- converted_weights_dir: `converted_weights/dreamx_world`
- source_layout: `raw_official`
- strict_load_status: `pass`
- conversion_script_status: `transformer_model_index_and_config_consistency_smoke_pass`
- model_index_status: `smoke_pass`
- passthrough_components: `Wan2.2 Diffusers scheduler, tokenizer, and text encoder are symlinked from official_weights/Wan2.2-TI2V-5B-Diffusers; VAE parity uses raw Wan2.2_VAE.pth with an explicit DreamX raw-to-FastVideo key mapper because official encode returns normalized latents.`
- retry_history: `none`
## Parity Commands
| Scope | Command | Last Result | Notes |
|---|---|---|---|
| transformer | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s` | strict_load_and_forward_parity_pass | 2026-07-01: converted real 5B-Cam transformer shards strict-load into dedicated `DreamXWorldTransformer3DModel` with 0 shape mismatches; official-vs-FastVideo small-input fp32 forward parity passes on CUDA (`diff_max=0.072533`, `diff_mean=0.008014`). |
| vae | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s` | encode_parity_pass | 2026-06-30: official DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE and encode parity passes after applying the same official latent normalization (`(mu - mean) / std`). |
| text_encoder | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s` | hidden_state_parity_pass | 2026-06-30: official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA using staged Wan2.2 text encoder/tokenizer weights and reference-only `xfuser` stubs. |
| scheduler | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s` | non_skip_pass | 2026-06-30: FastVideo FlowMatch scheduler matches official Diffusers timesteps and step output for DreamX default Flow sampler; `DreamXWorldPipeline` initializes FlowMatch with official `shift=3.0`. |
| camera_conditioning | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s` | non_skip_pass | 2026-06-29: 3 parameterized cases passed against official reference on CPU. |
| pipeline_config | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s` | pipeline_entry_preset_scheduler_modelinfo_and_camera_stage_smoke_pass | DreamX 5B-Cam PipelineConfig wires DiT/VAE/UMT5/Flow/TI2V settings and official `shift=3.0`; default preset is registered for `GD-ML/DreamX-World-5B-Cam`; local converted-style `model_index.json` resolves to `DreamXWorldPipeline`; the pipeline initializes FlowMatch, camera conditioning writes `batch.extra["dreamx_y_camera"]`, and generic denoising can pass it as `y_camera`. |
| ar_transformer | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_ar_conversion.py tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py -q -rs` | 6_passed_0_skipped | 2026-07-02: AR converter writes symlinked transformer layout/model_index; tiny official `CausalWanModel` vs FastVideo `DreamXWorldARTransformer3DModel` forward parity passes; real 5B `model.safetensors` strict-loads with zero missing/unexpected keys from `/tmp/converted_dreamx_world_ar`. |
| ar_pipeline_config | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -q -rs` | 10_passed_0_skipped | 2026-07-02: `DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, registry config selection for `GD-ML/DreamX-World-5B`, and `dreamx_world_5b_ar` preset pass. |
| ar_full_generation | `PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py` | generated_video_pass | 2026-07-02: A40 short full-generation smoke passed from `/tmp/converted_dreamx_world_ar` with 64x64, 9 frames, 4 denoise steps, `output_type=pil`, `save_video=True`; MP4 saved at `outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4` and decoded as 9 frames of `(64, 64, 3)` uint8. |
| ar_long_horizon | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py` | generated_video_pass | 2026-07-02: A40 long-horizon AR generation passed from `/tmp/converted_dreamx_world_ar` with 64x64, 1005 frames, 4 denoise steps, seed 4096; MP4 saved at `outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4` and decoded as 1005 frames of `(64, 64, 3)`; end-to-end generation latency was 231.68s after load. |
| ar_ssim_default | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download` | 1_passed_0_skipped | 2026-07-02: A40 default AR SSIM reference seeded locally at `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/`; helper removes stale generated base MP4 before generation; default params are 192x192, 81 frames, 4 steps, seed 2048, min SSIM 0.98. |
| ar_ssim_modal_l40s | `modal run /tmp/modal_dreamx_ar_ssim_git.py` | 1_passed_0_skipped | 2026-07-02: Modal L40S run checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded default `L40S_reference_videos` via `reference_videos_cli.py download`, used cached converted AR weights under `/root/data/dreamx_world_ar_converted`, seeded the missing AR L40S reference from generated output, reran a fresh generated-vs-reference compare successfully (`mean_ssim=1.0`), and exported the reference to Modal volume `hf-model-weights:dreamx_ar_ssim_l40s`. Downloaded local reference decodes as 81 frames of `(192, 192, 3)`. A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944. |
| pipeline_smoke | `python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs` | 4_passed_0_skipped | 2026-06-30: combined smoke/parity passed. 2026-07-01: smoke alone passed with real `image_path` TI2V coverage (`3 passed`), validating image load, TI2V preprocessing, VAE first-frame encode under CPU offload, camera conditioning, and 1-step latent generation from `converted_weights/dreamx_world`. Tests force `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA` to avoid the local FlashAttention-4 cute ABI mismatch. |
| basic_example | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA DREAMX_WORLD_MODEL_DIR=converted_weights/dreamx_world DREAMX_WORLD_IMAGE_PATH= DREAMX_WORLD_HEIGHT=64 DREAMX_WORLD_WIDTH=64 DREAMX_WORLD_NUM_FRAMES=9 DREAMX_WORLD_STEPS=1 DREAMX_WORLD_GUIDANCE=1.0 DREAMX_WORLD_OUTPUT_PATH=outputs_video/dreamx_world_example_smoke python examples/inference/basic/basic_dreamx_world.py` | generated_video_pass | 2026-06-30: example saved an MP4 under `outputs_video/dreamx_world_example_smoke`; imageio/ffmpeg decoded frame 0 as `(64, 64, 3)` uint8, fps 16, duration 0.56s. |
## Open Questions
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|---|---|---|---|---|---|
| Q001 | Should first PR expose only the `DreamX-World-5B-Cam` 5s camera-control mode and exclude AR long-horizon forcing? | user | prep | resolved | User approved starting with `DreamX-World-5B-Cam`; AR long-horizon is out of first-PR scope. |
| Q002 | Does FastVideo's existing Wan2.2 TI2V transformer support DreamX PRoPE/control adapter with a small extension, or is a DreamX-specific DiT required? | component:transformer | Phase 3 | resolved | Project guidance prefers a separate DreamX DiT for maintainability. DreamX PRoPE/control adapter now lives in `fastvideo/models/dits/dreamx_world.py`; Wan DiT/config have no DreamX-specific fields or `y_camera` signature. |
| Q003 | Which sampler is in first-PR scope: official default `Flow` only, or also `Flow_Unipc` and `Flow_DPM++`? | orchestrator | Phase 3 | resolved | First PR should support official default `Flow` only. FastVideo FlowMatch scheduler parity is non-skip PASS; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. |
| Q004 | Which HF token env var should be used if rate limits or gated Wan2.2 base weights require auth? | user | Phase 5 | resolved | No auth was required for the completed local downloads; keep using env var names only if future gated repos require auth. |
| Q005 | Should native FastVideo production code depend on DreamX reference-only packages such as `xfuser` or OpenCV? | user | Phase 3 | resolved | No. These packages may be used only for official reference/local parity setup; native FastVideo integration must remove that runtime requirement. |
| Q006 | Should AR handoff require a full generated long-horizon video in this no-HF-token/no-GPU-budget pass? | user/runtime | pipeline | resolved | A40 long-horizon generation passes with 1005 frames at 64x64/4 steps. AR default SSIM passes locally on A40 and on Modal L40S with a 192x192/81-frame reference. HF upload/publication remains a separate operation if the reference dataset should be updated upstream. |
| Q007 | Can A40 references stand in for L40S CI references? | quality | release | resolved | No. After deterministic AR noise fix, A40-vs-L40S reference mean SSIM is 0.7770, below the 0.98 same-device threshold. Publish the L40S-specific reference for CI. |
## Issues And Blockers
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|---|---|---|---|---|---|---|---|---|
| I001 | prep | official_env | medium | Official import initially failed because `xfuser` was missing. | `ModuleNotFoundError: No module named 'xfuser'` from `python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b"` | prep | resolved | Installed `xfuser==0.4.1`; import progressed. |
| I002 | prep | official_env | medium | Official import then failed because GUI OpenCV required missing system `libxcb.so.1`. | `ImportError: libxcb.so.1: cannot open shared object file` through `cv2` import in Diffusers ConsisID path. | prep | resolved | Installed `opencv-python-headless`; `import inference_dreamx5b` passed. |
| I003 | prep | weights | medium | HF repo has raw official transformer shards and no Diffusers `model_index.json`. | `inspect_hf_layout.py GD-ML/DreamX-World-5B-Cam --json` returned `source_layout=raw_official`, `needs_conversion=yes`, `model_index_class=null`. | conversion | resolved | Downloaded raw DreamX shards to `official_weights/dreamx_world`; converted transformer to `converted_weights/dreamx_world/transformer`; symlinked reusable Wan2.2 Diffusers components; real 5B transformer strict-load passes. |
| I004 | prep | dependencies | high | Official reference import required extra packages in the local environment, but FastVideo native runtime should not inherit those dependencies. | `xfuser==0.4.1` and `opencv-python-headless` were installed only to make `DreamX-World/inference_dreamx5b.py` import for reference/parity. | pipeline | resolved | Production DreamX FastVideo code uses native camera/image/video utilities and has no runtime `xfuser` or OpenCV import requirement; those packages remain reference-only local parity dependencies. |
| I005 | parity | transformer | medium | Transformer full forward parity initially failed in bf16 official harness. | Official CUDA bf16 LayerNorm path was unstable; fp32 small-input harness avoids that dtype issue and compares against FastVideo with single-process SP identity patches. | component:transformer | resolved | Official-vs-FastVideo forward parity now passes on CUDA with `diff_max=0.072533`, `diff_mean=0.008014`. |
| I006 | parity | vae/text_encoder | medium | VAE/text parity initially remained skipped after weights were staged. | Text official import needed a reference-only `xfuser` stub; VAE comparison initially used raw official normalized latents against FastVideo raw mu. | component:vae,component:text_encoder | resolved | Text hidden-state parity passes. VAE encode parity passes after raw key mapping and applying the official latent normalization to FastVideo output. |
| I007 | quality | pipeline_ti2v | medium | Real image-path TI2V smoke initially failed when `vae_cpu_offload=True` because DenoisingStage encoded the first frame while the VAE weights remained on CPU. | DreamX SSIM first run failed with `RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same` at `fastvideo/pipelines/stages/denoising.py` VAE encode. | pipeline | resolved | DenoisingStage now moves the VAE to `local_device` before TI2V first-frame encode; image-path pipeline smoke and DreamX SSIM both pass. |
| I008 | quality | ssim_helper | high | SSIM helper could compare against a stale generated base MP4 when a rerun saved the new video as `_1.mp4`. | Existing generated outputs made AR reference seeding appear to pass before a fresh generated-vs-reference compare. | quality | resolved | `run_text_to_video_similarity_test` and `run_image_to_video_similarity_test` now remove the stale generated base MP4 before generation. A40 and Modal L40S AR SSIM were rerun after the fix. |
| I009 | quality | ar_denoising | high | AR denoising added CUDA noise without using the request seed when the original generator was CPU-backed. | Fresh reruns against old AR references produced mean SSIM near 0.05. | pipeline | resolved | `DreamXWorldARCausalDenoisingStage` now derives a device-local generator from the request seed for AR noise. Fresh same-device A40 and L40S reruns pass with mean SSIM 1.0 after reseeding references. |
## Escape Hatches
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|---|---|---|---|---|---|---|
## Decisions
| Date | Decision | Rationale | Impact |
|---|---|---|---|
| 2026-06-29 | First PR scope is `DreamX-World-5B-Cam` only. | Cam mode is closest to existing Wan2.2 TI2V support; AR forcing needs separate causal/KV pipeline work. | Component inventory and parity focus on `inference_dreamx5b.py` and `pipeline_dreamxworld.py`. |
| 2026-06-29 | Do not install full DreamX requirements during prep. | Full requirements pin core FastVideo stack packages. | Installed only `xfuser==0.4.1` and `opencv-python-headless` to make official imports work. |
| 2026-06-29 | Treat HF DreamX-World-5B-Cam weights as raw official transformer layout requiring conversion. | HF inspection found no `model_index.json`. | Phase 5 must create `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py` after component prototype/key dumps. |
| 2026-06-29 | Do not add DreamX reference-only dependencies to FastVideo production requirements. | The current environment should remain the FastVideo environment; extra packages are only for official reference parity. | Native DreamX integration must avoid runtime `xfuser` and OpenCV requirements unless explicitly approved later. |
| 2026-06-29 | Implement DreamX camera conditioning as native FastVideo utility. | It is weightless and removes the need to import DreamX reference utilities at production runtime. | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` now has non-skip parity against official action-to-PRoPE tensors. |
| 2026-06-29 | First PR supports DreamX default `Flow` sampler only. | FastVideo FlowMatch Euler scheduler matches the official Diffusers scheduler for DreamX defaults. | Pipeline work can use FastVideo native FlowMatch scheduler; UniPC and DPM++ are out of first-PR scope. |
| 2026-07-01 | Keep DreamX PRoPE/control adapter in a dedicated DreamX DiT class. | Project guidance is that putting too much DreamX behavior into Wan makes the model hard to manage. | `fastvideo/models/dits/dreamx_world.py` defines `DreamXWorldTransformer3DModel`, `DreamXWorldTransformerBlock`, and `DreamXPropeSelfAttention`; `fastvideo/configs/models/dits/dreamx_world.py` owns DreamX adapter config fields; Wan DiT/config are unchanged from DreamX. |
| 2026-06-30 | Camera parity test loads official camera functions by file instead of importing the official `utils` package. | Official package initialization pulls unrelated dependencies that can require GUI OpenCV system libraries. | Camera parity remains non-skip without adding DreamX reference-only dependencies to FastVideo production requirements. |
| 2026-06-30 | Add DreamX-World-5B-Cam model and pipeline config helpers plus a conversion script. | Official HF DreamX 5B-Cam transformer config is 30 layers, hidden size 3072, 24 heads, 48 latent channels, plus Wan2.2 48-channel VAE and UMT5-XXL text encoder. | DreamX helpers wire DiT/VAE/UMT5/Flow/TI2V settings; `dreamx_world_to_diffusers.py` writes a FastVideo-loadable transformer config plus renamed safetensors; strict-load smoke passes on a tiny official DreamX transformer and the real 5B converted shards. |
| 2026-06-30 | Pass DreamX camera PRoPE condition through the FastVideo batch/denoising path. | DreamX transformer expects `y_camera={"viewmats", "K"}` at denoising time. | `DreamXWorldPipeline` is registered as a basic pipeline entry and initializes the official default FlowMatch scheduler; `dreamx_world_5b_cam` preset mirrors official 5B-Cam defaults; `DreamXWorldCameraConditioningStage` writes `batch.extra["dreamx_y_camera"]`; generic denoising filters and forwards it as `y_camera` only for compatible transformers. |
## Handoff Notes
- Prep, component parity, pipeline smoke/parity, and the basic example validation are complete for `DreamX-World-5B-Cam`.
- Official reference clone is staged at `DreamX-World/` and ignored by git.
- Workspace-local weights are staged: DreamX raw transformer shards under `official_weights/dreamx_world`, Wan2.2 raw base artifacts under `official_weights/Wan2.2-TI2V-5B`, and Wan2.2 Diffusers reusable components under `official_weights/Wan2.2-TI2V-5B-Diffusers`.
- Camera conditioning parity is active and passing without weights.
- Default Flow scheduler parity is active and passing without weights.
- Transformer has corrected official 5B-Cam architecture in dedicated DreamX DiT/config files, native PRoPE/control-adapter, conversion mapping, real converted 5B strict-load, and official-vs-FastVideo forward parity passing on CUDA. VAE encode parity and text hidden-state parity pass on CUDA. Pipeline entry/registry, local model_info resolution, preset, config, FlowMatch scheduler init, camera stage, denoising `y_camera` kwarg smokes, independent CUDA pipeline smoke/parity, and a small saved-video basic example pass.
- Full local DreamX component suite is non-skip PASS: `python -m pytest tests/local_tests/dreamx_world/ -q -rs` returned `26 passed` on 2026-07-01. Pipeline smoke/parity and SSIM quality regression are also non-skip PASS locally.
- Keep `xfuser` and OpenCV as reference-only parity dependencies. Do not add them
to FastVideo requirements or production imports.
## Quality Regression
- status: `added`
- test: `fastvideo/tests/ssim/test_dreamx_world_similarity.py`
- command: `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs`
- result: `1 passed, 0 skipped` on 2026-07-01
- reference: Local A40/TORCH_SDPA reference seeded from the generated candidate under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`. The test uses a deterministic generated input image, 64x64 request dimensions, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. Full-quality params are present for 480x832/161 frames/30 steps.
- note: Modal L40S seeding passed using the configured Modal profile and unauthenticated HF public downloads. HF upload/publication still requires `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` with write access; no token values were used or recorded.
## Final Handoff
```text
final_handoff:
prep_handoff_complete: yes
conversion_status: pass
components:
- name: transformer
reuse_or_port: ported_dedicated_dit
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: vae
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: text_encoder_tokenizer
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: scheduler
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: camera_conditioning
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: ar_transformer
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: ar_pipeline
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py plus AR smoke/SSIM commands listed above
parity_status: non_skip_pass
concerns_or_unknowns: none
pipeline_smoke: pass
pipeline_parity: pass
example_status: pass
quality_regression: added
local_tests_readme: tests/local_tests/dreamx_world/README.md
port_state_file: tests/local_tests/dreamx_world/PORT_STATUS.md
token_values_committed: no
runtime_third_party_model_imports: none
blockers: none
escape_hatch: none
```
| 2026-07-02 | Add DreamX-World-5B autoregressive support. | Official AR repo is raw single-safetensors layout and needs a dedicated causal/KV stage. | Added native AR DiT/config, identity converter, AR pipeline config/preset/registry, and targeted non-skip tests. Raw AR weights are staged at `/tmp/dreamx_world_ar_weights`; converted symlink layout at `/tmp/converted_dreamx_world_ar`. |
| 2026-07-02 | Validate DreamX-World-5B AR short full generation on A40. | Targeted parity/config tests prove components, but end-to-end runtime can still fail at scheduler, RoPE cache, device, or decode boundaries. | `DreamXWorldARPipeline` generated and saved a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`; the saved MP4 decodes to 9 frames. |
| 2026-07-02 | Validate DreamX-World-5B AR long-horizon and default SSIM on A40. | AR needs coverage beyond the 9-frame smoke to exercise longer KV/cache progression and a quality regression path. | 1005-frame 64x64/4-step generation passes and decodes; default AR SSIM test uses 192x192/81 frames because MS-SSIM requires short side >160. |
| 2026-07-02 | Validate DreamX-World-5B AR default SSIM on Modal L40S. | CI references are device-specific; A40 alone is not enough for L40S reference coverage. | Modal checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded existing default L40S references, seeded the missing AR reference, reran SSIM successfully (`mean_ssim=1.0`), and exported the L40S reference back to the workspace. |
+254
View File
@@ -0,0 +1,254 @@
# DreamX World Local Tests
Local-only parity and smoke tests for the `dreamx_world` FastVideo port. These
tests compare FastVideo against the official DreamX-World reference
implementation and are not expected to run in CI unless explicitly promoted
later.
Port progress, open questions, issues, and handoff notes live in
`tests/local_tests/dreamx_world/PORT_STATUS.md`.
## Reference Assets
| Field | Value |
|---|---|
| Model family | `dreamx_world` |
| First-PR scope | `DreamX-World-5B-Cam`; follow-up scope now includes `DreamX-World-5B` autoregressive forcing |
| Out-of-scope variants | none for the DreamX-World 5B/Cam paths currently ported |
| Workload types | I2V camera-control compatibility shim: image + prompt + action sequence to video |
| Official reference | `https://github.com/AMAP-ML/DreamX-World` |
| Local reference dir | `DreamX-World/` |
| Official commit/version | `221875811ba31f7eac6c3025b215c09ad2cefd1d` |
| HF weights | `GD-ML/DreamX-World-5B-Cam` |
| HF revision | default |
| Local weights dir | `official_weights/dreamx_world` |
| Source layout | `raw_official` |
| Needs conversion | `yes` |
Do not write token values in this file. Current token env var detected during
prep: `none`.
## Shared Environment Setup
Run from the FastVideo repo root in the same conda/env used for FastVideo. Do
not create a separate upstream environment for parity tests.
```bash
python ".agents/skills/add-model-01-prep/scripts/clone_reference_repo.py" \
"https://github.com/AMAP-ML/DreamX-World.git" \
"DreamX-World" \
--commit "221875811ba31f7eac6c3025b215c09ad2cefd1d" \
--update-gitignore
```
DreamX-World does not expose a packaging file for editable install. During prep
the official import check used `sys.path.insert(0, "DreamX-World")`.
Additional official deps installed into the current environment for imports:
```bash
uv pip install xfuser==0.4.1
uv pip install opencv-python-headless
```
These packages are for running the official DreamX reference during local
parity only. They must not become FastVideo production/runtime dependencies for
the native `dreamx_world` pipeline.
Do not install the full `DreamX-World/requirements.txt` without explicit
approval. It pins core FastVideo stack packages including `torch`, `torchvision`,
`triton`, `flash_attn`, and `diffusers`.
## Official Environment Status
```text
dependency_changes: installed official deps in current env
official_env_status: imports_ok
private_dep_stubs: none
blocked_on: none
```
Import check used during prep:
```bash
python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b; print('imports_ok')"
```
## Weight Setup
HF layout inspection found no root `model_index.json`; the repo contains
`config.json`, a safetensors index, and three transformer safetensors shards.
This is a raw official transformer layout and requires conversion before
FastVideo can load it through `VideoGenerator.from_pretrained`.
```bash
python ".agents/skills/add-model-01-prep/scripts/inspect_hf_layout.py" \
"GD-ML/DreamX-World-5B-Cam" \
--json
```
Weights have been staged workspace-locally. The raw DreamX transformer repo lives at `official_weights/dreamx_world`; Wan2.2 raw base artifacts live at `official_weights/Wan2.2-TI2V-5B`; Wan2.2 Diffusers reusable components live at `official_weights/Wan2.2-TI2V-5B-Diffusers`. To reproduce the DreamX download:
```bash
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
"GD-ML/DreamX-World-5B-Cam" \
"official_weights/dreamx_world"
```
### DreamX-World-5B Autoregressive Setup
The AR repository `GD-ML/DreamX-World-5B` is also raw official layout: no
`model_index.json`, root `config.json`, and a single `model.safetensors`. The
current environment has the raw AR checkpoint staged outside the workspace at
`/tmp/dreamx_world_ar_weights` to avoid workspace quota pressure. The converted
FastVideo layout is staged at `/tmp/converted_dreamx_world_ar` with the 21GB
transformer safetensors symlinked instead of copied.
```bash
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
"GD-ML/DreamX-World-5B" \
"/tmp/dreamx_world_ar_weights"
python scripts/checkpoint_conversion/dreamx_world_ar_to_diffusers.py \
--source /tmp/dreamx_world_ar_weights \
--output /tmp/converted_dreamx_world_ar \
--component-source official_weights/Wan2.2-TI2V-5B-Diffusers \
--symlink-components \
--symlink-transformer
```
AR production code uses `DreamXWorldARTransformer3DModel`,
`DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, and
`DreamXWorldARCausalDenoisingStage`. The AR DiT is a native FastVideo port of
the official Apache-2.0 `CausalWanModel`; it has no production DreamX, Diffusers
model-class, Transformers model-class, `xfuser`, or OpenCV import.
## Prototype And Conversion Artifacts
State-dict key/shape dumps are generated after FastVideo native prototypes exist
and are used to build the conversion mapping.
```text
official_key_dumps:
transformer: converted_weights/dreamx_world/_mapping/transformer_official_keys.json
fastvideo_key_dumps:
transformer: converted_weights/dreamx_world/_mapping/transformer_fastvideo_keys.json
conversion_script: scripts/checkpoint_conversion/dreamx_world_to_diffusers.py
conversion_script_status: transformer_model_index_and_config_consistency_smoke_pass
conversion_source_layout: raw_official
converted_weights_dir: converted_weights/dreamx_world
model_index_status: smoke_pass
strict_load_status: pass
```
The converter writes `transformer/` from raw DreamX shards. To create a full
FastVideo-loadable diffusers-style root, pass a Wan2.2 Diffusers directory as the
component source so reusable components are copied or symlinked before
`model_index.json` is emitted:
```bash
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
--source official_weights/dreamx_world \
--output converted_weights/dreamx_world \
--component-source /path/to/Wan2.2-TI2V-5B-Diffusers \
--symlink-components
```
## Expected Parity Tests
Planned local tests for this family:
| Component | Official files / args | Test | Concerns | Status |
|---|---|---|---|---|
| transformer | `DreamX-World/models/wan_transformer3d.py`; instantiated in `DreamX-World/inference_dreamx5b.py` with `cam_method=prope`, `add_control_adapter=True` | `tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py` | FastVideo uses dedicated `DreamXWorldTransformer3DModel`/`DreamXWorldConfig` files for DreamX 5B-Cam config, PRoPE, conversion mapping, real converted 5B strict-load PASS, and official-vs-FastVideo small-input forward parity PASS on CUDA. Wan DiT/config have no DreamX-specific adapter fields. | strict_load_and_forward_parity_pass |
| vae | `DreamX-World/models/wan_vae3_8.py`; `vae_type=AutoencoderKLWan3_8`, `vae_subpath=Wan2.2_VAE.pth` | `tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py` | DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE; encode parity passes after applying official latent normalization. | encode_parity_pass |
| text_encoder/tokenizer | `DreamX-World/models/wan_text_encoder.py`; T5 path from Wan2.2 base model | `tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py` | Official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA with staged Wan2.2 weights/tokenizer. | hidden_state_parity_pass |
| scheduler | Diffusers `FlowMatchEulerDiscreteScheduler`; selected by default `sampler_name=Flow` | `tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py` | First PR can support official default `Flow`; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. | non_skip_pass |
| camera_conditioning | `DreamX-World/utils/inference_utils.py`, `models/prope_utils.py`, `wan/modules/camera_prope.py` | `tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py` | Action sequence to PRoPE/control input must match official tensor shapes and values. | non_skip_pass |
| pipeline | `DreamX-World/pipeline/pipeline_dreamxworld.py`; call path in `DreamX-World/inference_dreamx5b.py` | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py` | DreamX PipelineConfig wires first-scope DiT/VAE/UMT5/Flow/TI2V settings, official `shift=3.0`, default preset values, FlowMatch scheduler initialization, and local `model_index.json` resolution; independent pipeline smoke covers real CUDA local load + latent generation, and parity compares public API output to worker-side explicit ForwardBatch execution. | pipeline_smoke_and_parity_pass |
| ar_transformer | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py`; instantiated by `DreamX-World/inference_ar_forcing.py` as `CausalWanModel` with `local_attn_size=12`, `sink_size=3`, `attn_compress=4` | `tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py` | Native `DreamXWorldARTransformer3DModel` keeps official identity key layout; tiny official-vs-FastVideo forward parity passes; real 5B AR safetensors strict-load passes from `/tmp/converted_dreamx_world_ar`. | tiny_forward_parity_and_real_strict_load_pass |
| ar_pipeline | `DreamX-World/pipeline/pipeline_causal_camera.py`; AR block/KV/context-noise loop | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; A40 short full-generation smoke | Dedicated `DreamXWorldARCausalDenoisingStage` implements blockwise KV forcing; raw HF repo must be converted before `VideoGenerator.from_pretrained` because it has no `model_index.json`; converted AR layout generated a 64x64/9-frame/4-step MP4 on A40. | config_registry_and_short_full_generation_pass |
Include reused components in parity. Reuse is accepted only after the FastVideo
component definition and official instantiation arguments have both been checked
and the component parity test passes non-skip.
Run the relevant tests with:
```bash
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_conversion.py -v -s
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
```
## Current Local Results
```bash
python -m pytest tests/local_tests/dreamx_world/ -v -s
# 2026-07-01: 26 passed, 0 skipped
# 2026-07-02: AR targeted suite passed: 16 passed, 0 skipped
PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py
# 2026-07-02: DreamX-World-5B AR short full-generation smoke passed on A40
# output: outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4
# decoded: 9 frames, (64, 64, 3), uint8
PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py
# 2026-07-02: DreamX-World-5B AR long-horizon generation passed on A40
# output: outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4
# decoded: 1005 frames, (64, 64, 3)
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download
# 2026-07-02: DreamX-World-5B AR default SSIM passed: 1 passed, 0 skipped
# local A40 reference: fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/
modal run /tmp/modal_dreamx_ar_ssim_git.py
# 2026-07-02: Modal L40S default SSIM passed: 1 passed, 0 skipped
# checked out post-fix dreamx-world-5b-cam branch commit
# first downloaded default L40S references with reference_videos_cli.py download
# seeded missing AR L40S reference, then reran a fresh generated-vs-reference compare
# Modal JSON: mean_ssim=1.0, min_ssim=1.0, max_ssim=1.0
# local L40S reference: fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/DreamX-World-5B/TORCH_SDPA/
# A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
# 2026-06-30: 4 passed, 0 skipped
# 2026-07-01: smoke image-path TI2V coverage passed separately with 3 passed, 0 skipped
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs
# 2026-07-01: 1 passed, 0 skipped
```
`camera_conditioning`, default `Flow` scheduler, DreamX component/pipeline configs, default preset, pipeline entry/registry, FlowMatch scheduler initialization, DreamX camera stage, denoising `y_camera` pass-through, conversion model-index/config-consistency checks, real converted 5B dedicated DreamX transformer strict-load and forward parity on CUDA, VAE encode parity on CUDA, text hidden-state parity on CUDA, native DreamX PRoPE branch structure smoke, AR transformer parity/strict-load, AR config/registry, and AR short full-generation smoke are non-skip PASS results. The full local DreamX component suite, independent pipeline smoke/parity suite, image-path TI2V smoke, and SSIM quality regression currently have zero skips. The basic 5B-Cam example was run against `converted_weights/dreamx_world` with a 64x64/9-frame/1-step saved-video smoke, and imageio decoded the generated MP4 successfully; the AR pipeline separately saved and decoded a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`.
## Review Notes
- Required before handoff: non-skip PASS for each required component parity
test, including reused components that own weights or numerical behavior.
- First PR scope originally targeted `DreamX-World-5B-Cam`; scope was later
expanded to include `DreamX-World-5B` AR support with a separate causal/KV
pipeline.
- AR support has targeted parity/config coverage, short full-generation smoke,
1005-frame A40 long-horizon generation, local A40 default SSIM coverage,
and Modal L40S default SSIM coverage. L40S validation was rerun after fixing
stale generated-output comparison in the SSIM helper and deterministic AR
noise seeding. HF reference publication remains a separate token-gated
release operation.
- FastVideo production code must not require `xfuser` or OpenCV just because the
official reference import needed them. Port camera/action preprocessing and
sequence-parallel behavior into existing FastVideo-native utilities or keep
reference-only imports inside local parity tests.
- Review agents should verify setup commands still match the PR, then run the
listed parity tests or report the exact blocker.
## Quality Regression
Quality regression is added in `fastvideo/tests/ssim/test_dreamx_world_similarity.py`. The 5B-Cam default test uses the workspace converted model root, `TORCH_SDPA`, a deterministic generated conditioning image, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. A local A40 reference was seeded under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`, and the test passed non-skip locally on 2026-07-01. The AR default test uses `/tmp/converted_dreamx_world_ar` or `/root/data/dreamx_world_ar_converted`, `TORCH_SDPA`, 192x192, 81 frames, 4 steps, seed 2048, and min SSIM 0.98; local A40 and Modal L40S references were seeded under `fastvideo/tests/ssim/reference_videos/default/{A40,L40S}_reference_videos/DreamX-World-5B/TORCH_SDPA/`, and both tests passed non-skip on 2026-07-02 after fresh generated-output cleanup was added to the SSIM helper. Modal L40S validation checked out `post-fix dreamx-world-5b-cam branch commit`, seeded the missing AR L40S reference, reran the test, and wrote `mean_ssim=1.0`. A cross-device A40-vs-L40S reference spot check produced mean SSIM 0.7770, so CI should use the L40S-specific reference rather than the A40 artifact. Full-quality params are present for 5B-Cam 480x832/161 frames/30 steps and AR 704x1280/1005 frames/4 steps; publishing references to the HF dataset remains a release operation requiring a write-capable HF token env var, never a raw token value.
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive conversion smoke tests."""
import json
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
from scripts.checkpoint_conversion.dreamx_world_ar_to_diffusers import (
MODEL_INDEX,
REUSED_COMPONENTS,
TRANSFORMER_CONFIG,
convert_transformer,
write_model_index,
)
def test_dreamx_world_ar_converter_writes_symlinked_transformer_and_model_index(tmp_path):
source = tmp_path / "raw"
source.mkdir()
raw_tensor = source / "model.safetensors"
raw_tensor.write_bytes(b"placeholder")
component_source = tmp_path / "wan22"
output = tmp_path / "dreamx_ar"
for component in REUSED_COMPONENTS:
component_dir = component_source / component
component_dir.mkdir(parents=True)
(component_dir / "config.json").write_text("{}\n")
convert_transformer(source, output, symlink_transformer=True)
write_model_index(output, component_source, symlink_components=True)
assert (output / "transformer" / "model.safetensors").is_symlink()
model_index = json.loads((output / "model_index.json").read_text())
assert model_index == MODEL_INDEX
assert model_index["_class_name"] == "DreamXWorldARPipeline"
assert model_index["transformer"] == ["diffusers", "DreamXWorldARTransformer3DModel"]
def test_dreamx_world_ar_transformer_config_matches_pipeline_dit_config():
dit_config = make_dreamx_world_5b_ar_dit_config()
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldARTransformer3DModel"
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
assert TRANSFORMER_CONFIG["local_attn_size"] == dit_config.local_attn_size
assert TRANSFORMER_CONFIG["sink_size"] == dit_config.sink_size
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
assert tuple(TRANSFORMER_CONFIG["cam_self_attn_layers"]) == dit_config.cam_self_attn_layers
def test_dreamx_world_ar_model_index_component_classes_are_registered():
for component in ("scheduler", "text_encoder", "transformer", "vae"):
class_name = MODEL_INDEX[component][1]
assert class_name in _LEGACY_FAST_VIDEO_MODELS
@@ -0,0 +1,182 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive transformer parity.
Coverage scope: both. The tiny forward parity compares FastVideo's native AR DiT
against the official DreamX ``CausalWanModel`` implementation with identical
weights. The real 5B checkpoint gate strict-loads the downloaded safetensors
through the FastVideo model class.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import pytest
import torch
from torch.testing import assert_close
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARArchConfig, DreamXWorldARConfig
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
from fastvideo.models.dits.dreamx_world_ar import DreamXWorldARTransformer3DModel
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
CONVERTED_AR_DIR = Path(os.getenv("DREAMX_WORLD_AR_CONVERTED_DIR", "/tmp/converted_dreamx_world_ar"))
CONVERTED_AR_HF_REPO = "FastVideo/DreamX-World-5B-Diffusers"
PARITY_SCOPE = "both"
def _tiny_config() -> DreamXWorldARConfig:
return DreamXWorldARConfig(
arch_config=DreamXWorldARArchConfig(
num_attention_heads=1,
attention_head_dim=8,
in_channels=4,
out_channels=4,
ffn_dim=16,
num_layers=1,
text_dim=8,
freq_dim=8,
text_len=4,
local_attn_size=2,
sink_size=1,
attn_compress=1,
cam_self_attn_layers=(0,),
))
def _load_official_tiny():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official DreamX reference missing: {OFFICIAL_REF_DIR}")
sys.path.insert(0, str(OFFICIAL_REF_DIR))
try:
from wan.modules import attention as official_attention
from wan.modules import causal_camera_model_2_2_prope_infinity as causal_module
from wan.modules import model_2_2 as official_model_2_2
from wan.modules.causal_camera_model_2_2_prope_infinity import CausalWanModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official AR transformer: {exc}")
official_attention.FLASH_ATTN_2_AVAILABLE = False
official_attention.FLASH_ATTN_3_AVAILABLE = False
def _sdpa_same_dtype(q, k, v, **kwargs):
del kwargs
out = torch.nn.functional.scaled_dot_product_attention(
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=0.0)
return out.transpose(1, 2).contiguous()
official_attention.attention = _sdpa_same_dtype
official_attention.flash_attention = _sdpa_same_dtype
official_model_2_2.flash_attention = _sdpa_same_dtype
causal_module.attention = _sdpa_same_dtype
return CausalWanModel(
model_type="ti2v",
patch_size=(1, 2, 2),
text_len=4,
in_dim=4,
dim=8,
ffn_dim=16,
freq_dim=8,
text_dim=8,
out_dim=4,
num_heads=1,
num_layers=1,
local_attn_size=2,
sink_size=1,
qk_norm=True,
cross_attn_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=(0,),
).eval()
def _make_inputs():
torch.manual_seed(123)
x = [torch.randn(4, 1, 4, 4)]
t = torch.zeros(1, 4, dtype=torch.long)
context = [torch.randn(2, 8)]
camera = {
"viewmats": torch.eye(4).reshape(1, 1, 4, 4).repeat(1, 4, 1, 1),
"K": torch.eye(3).reshape(1, 1, 3, 3).repeat(1, 4, 1, 1),
}
kv_cache = [{
"k": torch.zeros(1, 8, 1, 8),
"v": torch.zeros(1, 8, 1, 8),
"global_end_index": torch.tensor([0]),
"local_end_index": torch.tensor([0]),
"prope_k": torch.zeros(1, 8, 1, 8),
"prope_v": torch.zeros(1, 8, 1, 8),
"prope_global_end_index": torch.tensor([0]),
"prope_local_end_index": torch.tensor([0]),
}]
cross_cache = [{
"k": torch.zeros(1, 4, 1, 8),
"v": torch.zeros(1, 4, 1, 8),
"is_init": False,
}]
return x, t, context, camera, kv_cache, cross_cache
def test_dreamx_world_ar_tiny_forward_matches_official():
official = _load_official_tiny()
# The official init_weights zero-inits the output head (head.head.weight and
# biases), so both models would output exactly zero and the comparison would
# pass vacuously. Randomize the head (deterministically) before copying the
# state dict so the outputs reflect the internal computation.
generator = torch.Generator().manual_seed(7)
with torch.no_grad():
official.head.head.weight.normal_(std=0.5, generator=generator)
official.head.head.bias.normal_(std=0.5, generator=generator)
fastvideo = DreamXWorldARTransformer3DModel(_tiny_config(), {}).eval()
fastvideo.load_state_dict(official.state_dict(), strict=True)
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
official_out = official(x=x, t=t, context=context, seq_len=4, y_camera=camera,
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
fastvideo_out = fastvideo(x=x, t=t, context=context, seq_len=4, y_camera=camera,
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
assert official_out.abs().max() > 0, "official output is all-zero; parity comparison is vacuous"
assert_close(fastvideo_out, official_out, atol=1e-5, rtol=1e-5)
def test_dreamx_world_ar_5b_config_matches_official_shape():
config = make_dreamx_world_5b_ar_dit_config()
assert config.num_layers == 30
assert config.num_attention_heads == 24
assert config.attention_head_dim == 128
assert config.hidden_size == 3072
assert config.ffn_dim == 14336
assert config.local_attn_size == 12
assert config.sink_size == 3
assert config.attn_compress == 4
assert config.cam_self_attn_layers == tuple(range(30))
def test_dreamx_world_ar_converted_5b_transformer_strict_loads():
transformer_dir = CONVERTED_AR_DIR / "transformer"
if not transformer_dir.exists():
# No local conversion: pull the published Diffusers transformer from the hub.
from huggingface_hub import snapshot_download
transformer_dir = Path(snapshot_download(CONVERTED_AR_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
with torch.device("meta"):
model = DreamXWorldARTransformer3DModel(make_dreamx_world_5b_ar_dit_config(), {})
incompatible = load_model_from_full_model_state_dict(
model,
safetensors_weights_iterator(resolve_safetensors_files(str(transformer_dir)), to_cpu=True),
device=torch.device("cpu"),
param_dtype=torch.bfloat16,
strict=True,
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in model.parameters())
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World camera-conditioning parity against the official reference.
Coverage scope: implementation_subcomponent. This verifies the weightless
action-sequence to PRoPE camera tensor path used by DreamX-World-5B-Cam.
"""
from __future__ import annotations
import os
from pathlib import Path
import pytest
import torch
from torch.testing import assert_close
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
PARITY_SCOPE = "implementation_subcomponent"
def _load_official_functions():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
try:
import importlib.util
pose_path = OFFICIAL_REF_DIR / "utils" / "pose_utils.py"
pose_spec = importlib.util.spec_from_file_location(
"dreamx_world_pose_utils", pose_path)
if pose_spec is None or pose_spec.loader is None:
raise RuntimeError(f"Cannot load DreamX pose_utils: {pose_path}")
pose_module = importlib.util.module_from_spec(pose_spec)
pose_spec.loader.exec_module(pose_module)
source = (OFFICIAL_REF_DIR / "utils" / "inference_utils.py").read_text()
source = source.replace(
"from .pose_utils import interpolate_camera_poses\n", "")
namespace = {"interpolate_camera_poses": pose_module.interpolate_camera_poses}
exec(compile(source, str(OFFICIAL_REF_DIR / "utils" / "inference_utils.py"), "exec"), namespace)
except Exception as exc: # noqa: BLE001 - local parity should skip missing reference deps.
pytest.skip(f"Cannot load DreamX camera reference: {exc}")
return namespace["ActionToPoseFromID"], namespace["GetPoseEmbedsFromPosesPrope"]
def _official_camera_condition(
action_seq: list[str],
action_speed_list: list[float],
*,
num_frames: int,
height: int,
width: int,
dtype: torch.dtype,
):
action_to_pose, get_pose_embeds = _load_official_functions()
duration = -(-num_frames // len(action_seq))
poses = action_to_pose(action_seq, action_speed_list, duration=duration)[:num_frames]
condition, _ = get_pose_embeds(poses, height, width, len(poses), False, 0, dtype=dtype, device="cpu")
return condition
@pytest.mark.parametrize(
("action_seq", "action_speed_list", "num_frames"),
[
(["w"], [4], 81),
(["wj", "d"], [4, 6], 121),
(["i", "k", "l"], [3, 5, 2], 85),
],
)
def test_dreamx_world_camera_conditioning_matches_official(action_seq, action_speed_list, num_frames):
dtype = torch.float32
official = _official_camera_condition(
action_seq,
action_speed_list,
num_frames=num_frames,
height=704,
width=1280,
dtype=dtype,
)
fastvideo = build_dreamx_camera_condition(
action_seq,
action_speed_list,
num_frames=num_frames,
height=704,
width=1280,
dtype=dtype,
device="cpu",
)
assert official.keys() == fastvideo.keys() == {"viewmats", "K"}
for key in ("viewmats", "K"):
assert official[key].shape == fastvideo[key].shape
diff = (official[key] - fastvideo[key]).abs()
print(f"{key}: shape={tuple(fastvideo[key].shape)} diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
assert_close(fastvideo[key], official[key], atol=1e-5, rtol=1e-5)
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World conversion script smoke tests."""
import json
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import (
MODEL_INDEX,
REUSED_COMPONENTS,
TRANSFORMER_CONFIG,
_copy_or_link_component,
write_model_index,
)
def test_dreamx_world_converter_writes_full_model_index_with_reused_components(tmp_path):
component_source = tmp_path / "wan22"
output = tmp_path / "dreamx"
output.mkdir()
for component in REUSED_COMPONENTS:
component_dir = component_source / component
component_dir.mkdir(parents=True)
(component_dir / "config.json").write_text("{}\n")
write_model_index(output, component_source, symlink_components=True)
model_index = json.loads((output / "model_index.json").read_text())
assert model_index == MODEL_INDEX
assert model_index["_class_name"] == "DreamXWorldPipeline"
assert model_index["transformer"] == ["diffusers", "DreamXWorldTransformer3DModel"]
for component in REUSED_COMPONENTS:
assert (output / component).is_symlink()
def test_dreamx_world_transformer_config_has_camera_adapter_enabled():
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldTransformer3DModel"
assert TRANSFORMER_CONFIG["add_control_adapter"] is True
assert TRANSFORMER_CONFIG["cam_method"] == "prope"
assert TRANSFORMER_CONFIG["num_layers"] == 30
def test_dreamx_world_converter_transformer_config_matches_pipeline_dit_config():
dit_config = make_dreamx_world_5b_cam_dit_config()
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
assert TRANSFORMER_CONFIG["cross_attn_norm"] == dit_config.cross_attn_norm
assert TRANSFORMER_CONFIG["qk_norm"] == dit_config.qk_norm
assert TRANSFORMER_CONFIG["add_control_adapter"] == dit_config.add_control_adapter
assert TRANSFORMER_CONFIG["cam_method"] == dit_config.cam_method
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
assert TRANSFORMER_CONFIG["cam_self_attn_layers"] == dit_config.cam_self_attn_layers
def test_dreamx_world_model_index_component_classes_are_registered():
for component in ("scheduler", "text_encoder", "transformer", "vae"):
class_name = MODEL_INDEX[component][1]
assert class_name in _LEGACY_FAST_VIDEO_MODELS
def test_dreamx_world_copy_or_link_component_keeps_broken_symlink(tmp_path):
component_source = tmp_path / "wan22"
src = component_source / "scheduler"
src.mkdir(parents=True)
output = tmp_path / "dreamx"
output.mkdir()
dst = output / "scheduler"
dst.symlink_to(tmp_path / "missing_scheduler", target_is_directory=True)
_copy_or_link_component("scheduler", component_source, output, symlink=True)
assert dst.is_symlink()
@@ -0,0 +1,306 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline config and conditioning smoke tests."""
from types import SimpleNamespace
import json
import numpy as np
import torch
from fastvideo.api.presets import get_preset
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.pipeline_registry import PipelineType, import_pipeline_classes
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
DreamXCamera,
_interpolate_camera_poses,
build_dreamx_camera_condition,
)
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.registry import get_default_preset, get_model_info, get_pipeline_config_cls_from_name
def test_dreamx_world_5b_cam_pipeline_config_wires_first_scope_components():
config = DreamXWorld5BCamPipelineConfig()
assert config.flow_shift == 3.0
assert config.ti2v_task is True
assert config.expand_timesteps is True
assert config.dit_config.expand_timesteps is True
assert config.dit_config.num_layers == 30
assert config.dit_config.add_control_adapter is True
assert config.dit_config.cam_method == "prope"
assert config.vae_config.load_encoder is True
assert config.vae_config.load_decoder is True
assert config.vae_config.z_dim == 48
assert config.vae_config.scale_factor_temporal == 4
assert config.vae_config.scale_factor_spatial == 16
assert len(config.text_encoder_configs) == 1
text_config = config.text_encoder_configs[0]
assert text_config.prefix == "umt5"
assert text_config.vocab_size == 256384
assert text_config.d_model == 4096
assert config.text_encoder_precisions == ("bf16",)
def test_dreamx_world_pipeline_registry_discovers_entrypoint():
pipelines = import_pipeline_classes(PipelineType.BASIC)
assert pipelines["basic"]["DreamXWorldPipeline"] is DreamXWorldPipeline
def test_dreamx_world_local_model_index_resolves_model_info(tmp_path):
model_dir = tmp_path / "DreamX-World-5B-Cam-converted"
model_dir.mkdir()
for component in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
(model_dir / component).mkdir()
(model_dir / "model_index.json").write_text(json.dumps({
"_class_name": "DreamXWorldPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}) + "\n")
info = get_model_info(str(model_dir), pipeline_type=PipelineType.BASIC, workload_type=WorkloadType.I2V)
assert info.pipeline_cls is DreamXWorldPipeline
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
def test_dreamx_world_model_path_resolves_pipeline_config():
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B-Cam") is DreamXWorld5BCamPipelineConfig
def test_dreamx_world_default_preset_is_registered():
preset_name = get_default_preset("GD-ML/DreamX-World-5B-Cam")
preset = get_preset(preset_name, "dreamx_world")
assert preset.name == "dreamx_world_5b_cam"
assert preset.workload_type == "i2v"
assert preset.defaults["height"] == 480
assert preset.defaults["width"] == 832
assert preset.defaults["num_frames"] == 161
assert preset.defaults["num_inference_steps"] == 30
assert preset.defaults["guidance_scale"] == 5.0
def test_dreamx_world_pipeline_initializes_official_flow_scheduler():
pipeline = DreamXWorldPipeline.__new__(DreamXWorldPipeline)
pipeline.modules = {}
fastvideo_args = SimpleNamespace(pipeline_config=DreamXWorld5BCamPipelineConfig())
pipeline.initialize_pipeline(fastvideo_args)
scheduler = pipeline.modules["scheduler"]
assert isinstance(scheduler, FlowMatchEulerDiscreteScheduler)
assert scheduler.config.shift == 3.0
def test_dreamx_world_camera_conditioning_stage_sets_y_camera_extra():
batch = ForwardBatch(
data_type="t2v",
action_list=["wj", "d"],
action_speed_list=[4, 6],
num_frames=17,
height=704,
width=1280,
latents=torch.zeros(1, 16, 5, 44, 80),
)
stage = DreamXWorldCameraConditioningStage()
out = stage.forward(batch, fastvideo_args=object())
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
expected = build_dreamx_camera_condition(
["wj", "d"],
[4, 6],
num_frames=17,
height=704,
width=1280,
dtype=torch.float32,
device="cpu",
)
assert set(y_camera) == {"viewmats", "K"}
for key, expected_value in expected.items():
assert y_camera[key].shape == (1, *expected_value.shape)
torch.testing.assert_close(y_camera[key][0], expected_value)
assert stage.verify_output(out, fastvideo_args=object()).is_valid()
def test_dreamx_world_denoising_kwargs_filter_for_y_camera():
y_camera = {"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)}
stage = DenoisingStage.__new__(DenoisingStage)
def accepts_y_camera(hidden_states, encoder_hidden_states, timestep, y_camera=None):
return y_camera
def no_y_camera(hidden_states, encoder_hidden_states, timestep):
return hidden_states
assert stage.prepare_extra_func_kwargs(accepts_y_camera, {"y_camera": y_camera}) == {"y_camera": y_camera}
assert stage.prepare_extra_func_kwargs(no_y_camera, {"y_camera": y_camera}) == {}
def test_dreamx_world_ar_pipeline_config_wires_components():
config = DreamXWorld5BARPipelineConfig()
assert config.is_causal is True
assert config.flow_shift == 5.0
assert config.dmd_denoising_steps == (1000, 750, 500, 250)
assert config.warp_denoising_step is True
assert config.context_noise == 0.1
assert config.dit_config.arch_config.local_attn_size == 12
assert config.dit_config.arch_config.sink_size == 3
assert config.dit_config.arch_config.attn_compress == 4
def test_dreamx_world_ar_pipeline_registry_and_preset():
from fastvideo.api.presets import get_preset
from fastvideo.registry import get_pipeline_config_cls_from_name, get_preset_selection
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B") is DreamXWorld5BARPipelineConfig
preset_name, family = get_preset_selection("GD-ML/DreamX-World-5B")
assert (preset_name, family) == ("dreamx_world_5b_ar", "dreamx_world")
preset = get_preset("dreamx_world_5b_ar", "dreamx_world")
assert preset.defaults["num_inference_steps"] == 4
assert DreamXWorldARPipeline.pipeline_config_cls is DreamXWorld5BARPipelineConfig
def test_dreamx_world_camera_conditioning_stage_expands_scalar_speed():
batch = ForwardBatch(
data_type="t2v",
action_list=["w", "d"],
action_speed_list=2.0,
num_frames=17,
height=704,
width=1280,
latents=torch.zeros(1, 16, 5, 44, 80),
)
stage = DreamXWorldCameraConditioningStage()
out = stage.forward(batch, fastvideo_args=object())
assert set(out.extra[DREAMX_Y_CAMERA_KEY]) == {"viewmats", "K"}
def test_dreamx_world_camera_interpolation_handles_single_camera():
camera = DreamXCamera(
fx=0.8,
fy=0.8,
cx=0.5,
cy=0.5,
w2c_mat=np.eye(4, dtype=np.float64),
)
out = _interpolate_camera_poses(
[camera],
src_indices=np.array([0.0]),
tgt_indices=np.array([0.0, 1.0, 2.0]),
)
assert out == [camera, camera, camera]
def test_dreamx_world_ar_cache_initializes_camera_self_attention_entries():
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
transformer = SimpleNamespace(
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn), SimpleNamespace(cam_self_attn=cam_self_attn)],
num_attention_heads=2,
attention_head_dim=4,
)
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
stage.transformer = transformer
stage.num_transformer_blocks = 2
stage.local_attn_size = 6
caches = stage._initialize_kv_cache(
batch_size=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_seq_length=7,
)
assert len(caches) == 2
assert caches[0]["k"].shape == (1, 42, 2, 4)
assert caches[0]["prope_k"].shape == (1, 42, 3, 5)
assert caches[0]["prope_v"].shape == (1, 42, 3, 5)
assert int(caches[0]["prope_global_end_index"].item()) == 0
assert int(caches[0]["prope_local_end_index"].item()) == 0
def test_dreamx_world_ar_context_noise_fraction_maps_to_scheduler_timestep():
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(0.1) == 100
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(100) == 100
def test_dreamx_world_ar_context_update_advances_camera_cache_indices():
class DummyTransformer:
def __call__(self, *, hidden_states, encoder_hidden_states, timestep, y_camera, kv_cache, crossattn_cache,
current_start):
del encoder_hidden_states, y_camera, crossattn_cache
assert current_start == 0
assert timestep.unique().tolist() == [100]
new_tokens = timestep.shape[1]
for cache in kv_cache:
cache["local_end_index"] += new_tokens
cache["global_end_index"] += new_tokens
cache["prope_local_end_index"] += new_tokens
cache["prope_global_end_index"] += new_tokens
cache["k"][:, :new_tokens] = 1
cache["prope_k"][:, :new_tokens] = 1
return hidden_states
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
cache_transformer = SimpleNamespace(
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn)],
num_attention_heads=2,
attention_head_dim=4,
)
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
stage.transformer = cache_transformer
stage.num_transformer_blocks = 1
stage.local_attn_size = 6
caches = stage._initialize_kv_cache(
batch_size=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_seq_length=2,
)
# Keep the cache allocation source separate from the callable transformer used by _update_context_cache.
stage.transformer = DummyTransformer()
stage._update_context_cache(
block_latents=torch.zeros(1, 4, 3, 2, 2),
context=[torch.zeros(2, 4)],
camera_block={"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)},
kv_cache=caches,
crossattn_cache=[{}],
start=0,
frame_seq_length=2,
target_dtype=torch.float32,
autocast_enabled=False,
context_noise=0.1,
)
assert int(caches[0]["local_end_index"].item()) == 6
assert int(caches[0]["prope_local_end_index"].item()) == 6
assert caches[0]["k"][:, :6].sum().item() == 48
assert caches[0]["prope_k"][:, :6].sum().item() == 90
@@ -0,0 +1,53 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World default Flow scheduler parity.
Coverage scope: implementation_subcomponent. DreamX-World-5B-Cam defaults to
Diffusers FlowMatchEulerDiscreteScheduler for sampler_name=Flow. This test
checks that FastVideo's native FlowMatchEulerDiscreteScheduler matches the
timestep schedule and Euler step used by the official default path.
"""
from __future__ import annotations
from pathlib import Path
import inspect
import torch
from diffusers import FlowMatchEulerDiscreteScheduler as OfficialFlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler as FastVideoFlowMatchEulerDiscreteScheduler,
)
REPO_ROOT = Path(__file__).resolve().parents[3]
PARITY_SCOPE = "implementation_subcomponent"
def _scheduler_kwargs(cls):
config_path = REPO_ROOT / "DreamX-World" / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
config = OmegaConf.load(config_path)
raw_kwargs = OmegaConf.to_container(config["scheduler_kwargs"])
signature = inspect.signature(cls)
return {key: value for key, value in raw_kwargs.items() if key in signature.parameters}
def test_dreamx_world_flow_scheduler_timesteps_and_step_match():
official = OfficialFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(OfficialFlowMatchEulerDiscreteScheduler))
fastvideo = FastVideoFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(FastVideoFlowMatchEulerDiscreteScheduler))
official.set_timesteps(50, device="cpu", mu=1)
fastvideo.set_timesteps(50, device="cpu", mu=1)
assert_close(fastvideo.timesteps, official.timesteps, atol=0, rtol=0)
assert_close(fastvideo.sigmas, official.sigmas, atol=0, rtol=0)
torch.manual_seed(7)
sample = torch.randn(1, 4, 2, 8, 8)
model_output = torch.randn_like(sample)
timestep = official.timesteps[3]
official_prev = official.step(model_output, timestep, sample, return_dict=False)[0]
fastvideo_prev = fastvideo.step(model_output, fastvideo.timesteps[3], sample, return_dict=False)[0]
diff = (official_prev - fastvideo_prev).abs()
print(f"scheduler diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
assert_close(fastvideo_prev, official_prev, atol=0, rtol=0)
@@ -0,0 +1,159 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World Wan T5 encoder reuse parity scaffold.
Coverage scope: implementation_subcomponent. It records the official
WanT5EncoderModel loading path and FastVideo T5 target for later activation
with staged Wan2.2 base text encoder/tokenizer weights.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from transformers import AutoTokenizer
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_cam_text_encoder_config,
)
from fastvideo.models.loader.component_loader import TextEncoderLoader
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
WAN_DIFFUSERS_DIR = Path(os.getenv("DREAMX_WORLD_WAN_DIFFUSERS_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B-Diffusers"))
PARITY_SCOPE = "implementation_subcomponent"
def _add_official_to_path():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _text_kwargs():
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
return OmegaConf.to_container(config["text_encoder_kwargs"])
def _patch_single_process_text_parallel(monkeypatch):
import fastvideo.layers.linear as fastvideo_linear
import fastvideo.layers.vocab_parallel_embedding as fastvideo_embedding
import fastvideo.models.encoders.t5 as fastvideo_t5
for module in (fastvideo_t5, fastvideo_embedding, fastvideo_linear):
if hasattr(module, "get_tp_rank"):
monkeypatch.setattr(module, "get_tp_rank", lambda: 0)
if hasattr(module, "get_tp_world_size"):
monkeypatch.setattr(module, "get_tp_world_size", lambda: 1)
monkeypatch.setattr(fastvideo_embedding, "tensor_model_parallel_all_reduce", lambda x: x)
def _load_official_text_encoder(device, dtype):
_add_official_to_path()
_install_xfuser_stub()
text_path = WAN_BASE_DIR / "models_t5_umt5-xxl-enc-bf16.pth"
if not text_path.exists():
pytest.skip(f"Wan2.2 text encoder weights missing: {text_path}")
try:
from models import WanT5EncoderModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX text encoder: {exc}")
model = WanT5EncoderModel.from_pretrained(
str(text_path), additional_kwargs=_text_kwargs(), low_cpu_mem_usage=True, torch_dtype=dtype
)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_text_encoder(device, dtype, monkeypatch):
text_encoder_path = WAN_DIFFUSERS_DIR / "text_encoder"
if not text_encoder_path.exists():
pytest.skip(f"Wan2.2 Diffusers text encoder missing: {text_encoder_path}")
_patch_single_process_text_parallel(monkeypatch)
pipeline_config = DreamXWorld5BCamPipelineConfig()
pipeline_config.text_encoder_configs[0]._fsdp_shard_conditions = []
args = FastVideoArgs(
model_path=str(text_encoder_path),
pipeline_config=pipeline_config,
text_encoder_cpu_offload=(device.type == "cpu"),
)
args.model_paths = {}
return TextEncoderLoader().load(str(text_encoder_path), args).to(device=device, dtype=dtype).eval()
def test_dreamx_world_text_encoder_config_matches_umt5_xxl_shape():
config = make_dreamx_world_5b_cam_text_encoder_config()
assert config.vocab_size == 256384
assert config.d_model == 4096
assert config.d_kv == 64
assert config.d_ff == 10240
assert config.num_heads == 64
assert config.num_layers == 24
assert config.relative_attention_num_buckets == 32
assert config.dropout_rate == 0.0
assert config.text_len == 512
assert config.prefix == "umt5"
def test_dreamx_world_fastvideo_text_encoder_loads_staged_weights(monkeypatch):
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
model = _load_fastvideo_text_encoder(device, torch.bfloat16, monkeypatch)
assert model.__class__.__name__ == "UMT5EncoderModel"
assert next(model.parameters()).device.type == device.type
assert next(model.parameters()).dtype == torch.bfloat16
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for text encoder parity.")
def test_dreamx_world_text_encoder_parity_scaffold(monkeypatch):
device = torch.device("cuda:0")
dtype = torch.bfloat16
official = _load_official_text_encoder(device, dtype)
fastvideo = _load_fastvideo_text_encoder(device, dtype, monkeypatch)
tokenizer_path = WAN_BASE_DIR / "google" / "umt5-xxl"
if not tokenizer_path.exists():
pytest.skip(f"Wan2.2 tokenizer missing: {tokenizer_path}")
tokenizer = AutoTokenizer.from_pretrained(str(tokenizer_path))
batch = tokenizer(["A quiet forest trail at sunrise."], padding="max_length", max_length=512, return_tensors="pt")
input_ids = batch.input_ids.to(device)
attention_mask = batch.attention_mask.to(device)
with torch.inference_mode():
official_hidden = official(input_ids, attention_mask=attention_mask)[0].float().cpu()
fastvideo_hidden = fastvideo(input_ids, attention_mask=attention_mask).last_hidden_state.float().cpu()
assert official_hidden.shape == fastvideo_hidden.shape
assert_close(fastvideo_hidden, official_hidden, atol=1e-3, rtol=1e-3)
@@ -0,0 +1,335 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World transformer parity scaffold.
Coverage scope: both. The official side loads DreamX-World-5B-Cam through
Wan2_2Transformer3DModel.from_pretrained with PRoPE camera control enabled.
The FastVideo side strict-loads the converted DreamX transformer weights into
the native DreamX-World DiT implementation.
"""
from __future__ import annotations
import math
import os
from pathlib import Path
import sys
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.forward_context import set_forward_context
from fastvideo.configs.models.dits.dreamx_world import (
DreamXWorldArchConfig, DreamXWorldConfig)
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.models.dits.dreamx_world import (
DreamXPropeSelfAttention, DreamXWorldTransformer3DModel,
DreamXWorldTransformerBlock)
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import map_transformer_key
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
LOCAL_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "dreamx_world"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
CONVERTED_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "dreamx_world"))
CONVERTED_HF_REPO = "FastVideo/DreamX-World-5B-Cam-Diffusers"
PARITY_SCOPE = "both"
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _make_tiny_dreamx_config() -> DreamXWorldConfig:
return DreamXWorldConfig(
arch_config=DreamXWorldArchConfig(
num_attention_heads=1,
attention_head_dim=8,
in_channels=16,
out_channels=16,
ffn_dim=32,
num_layers=1,
cross_attn_norm=True,
qk_norm="rms_norm_across_heads",
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=None,
))
def _add_official_to_path() -> None:
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _official_transformer_kwargs() -> dict:
config_path = OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
if not config_path.exists():
pytest.skip(f"DreamX Wan config missing: {config_path}")
config = OmegaConf.load(config_path)
kwargs = OmegaConf.to_container(config["transformer_additional_kwargs"])
kwargs["cam_method"] = "prope"
kwargs["add_control_adapter"] = True
return kwargs
def _load_official_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
_add_official_to_path()
_install_xfuser_stub()
if not LOCAL_WEIGHTS_DIR.exists():
pytest.skip(f"DreamX transformer weights missing: {LOCAL_WEIGHTS_DIR}")
try:
from models import Wan2_2Transformer3DModel
except Exception as exc: # noqa: BLE001 - local parity should skip missing refs.
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
model = Wan2_2Transformer3DModel.from_pretrained(
str(LOCAL_WEIGHTS_DIR),
transformer_additional_kwargs=_official_transformer_kwargs(),
torch_dtype=dtype,
)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
model = _load_fastvideo_transformer_strict(torch.device("cpu"), dtype)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_transformer_strict(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
transformer_dir = CONVERTED_WEIGHTS_DIR / "transformer"
if not transformer_dir.exists():
# No local conversion: pull the published Diffusers transformer from the hub.
from huggingface_hub import snapshot_download
transformer_dir = Path(snapshot_download(CONVERTED_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
safetensors_files = resolve_safetensors_files(str(transformer_dir))
config = make_dreamx_world_5b_cam_dit_config()
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
original_get_sp_world_size = fastvideo_dreamx.get_sp_world_size
fastvideo_dreamx.get_sp_world_size = lambda: 1
try:
with torch.device("meta"):
model = DreamXWorldTransformer3DModel(config=config, hf_config={})
finally:
fastvideo_dreamx.get_sp_world_size = original_get_sp_world_size
incompatible = load_model_from_full_model_state_dict(
model,
safetensors_weights_iterator(safetensors_files, to_cpu=True),
device=device,
param_dtype=dtype,
strict=True,
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in model.parameters())
return model.to(device=device, dtype=dtype).eval()
def _make_inputs(device: torch.device, dtype: torch.dtype):
torch.manual_seed(1234)
num_frames = 5
height = 64
width = 64
latent_frames = (num_frames - 1) // 4 + 1
latent_h = height // 16
latent_w = width // 16
x = torch.randn(1, 48, latent_frames, latent_h, latent_w, device=device, dtype=dtype)
context = [torch.randn(16, 4096, device=device, dtype=dtype)]
seq_len = math.ceil((latent_h * latent_w) / 4 * latent_frames)
timestep = torch.full((1, seq_len), 250, device=device, dtype=torch.long)
camera = build_dreamx_camera_condition(
["w"], [4], num_frames=num_frames, height=height, width=width, dtype=dtype, device=device
)
camera = {key: value.unsqueeze(0) for key, value in camera.items()}
return {"x": [x[0]], "context": context, "t": timestep, "seq_len": seq_len, "y_camera": camera}
def _run_official(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
with torch.inference_mode():
output = model(**inputs)
if isinstance(output, list):
output = torch.stack(output, dim=0)
assert torch.is_tensor(output), f"official output is not a tensor: {type(output)}"
return output.detach().float().cpu()
def _run_fastvideo(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
hidden_states = torch.stack(inputs["x"], dim=0)
encoder_hidden_states = torch.stack([
torch.cat([inputs["context"][0], inputs["context"][0].new_zeros(512 - inputs["context"][0].shape[0], 4096)])
])
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
output = model(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=inputs["t"],
y_camera=inputs["y_camera"],
)
assert torch.is_tensor(output), f"FastVideo output is not a tensor: {type(output)}"
return output.detach().float().cpu()
def test_dreamx_world_conversion_mapping_strict_load_smoke(monkeypatch):
_add_official_to_path()
_install_xfuser_stub()
try:
from models import Wan2_2Transformer3DModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
official = Wan2_2Transformer3DModel(
dim=8,
ffn_dim=32,
num_heads=1,
num_layers=1,
add_control_adapter=True,
cam_method="prope",
)
official_state = official.state_dict()
diffusers_like_state = {
map_transformer_key(key): value.detach().clone()
for key, value in official_state.items()
}
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
with torch.device("meta"):
fastvideo = DreamXWorldTransformer3DModel(config=_make_tiny_dreamx_config(), hf_config={})
incompatible = load_model_from_full_model_state_dict(
fastvideo,
iter(diffusers_like_state.items()),
device=torch.device("cpu"),
param_dtype=torch.float32,
strict=True,
param_names_mapping=get_param_names_mapping(fastvideo.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in fastvideo.parameters())
def test_dreamx_world_5b_cam_dit_config_matches_official_shape():
config = make_dreamx_world_5b_cam_dit_config()
assert config.num_layers == 30
assert config.num_attention_heads == 24
assert config.attention_head_dim == 128
assert config.hidden_size == 3072
assert config.ffn_dim == 14336
assert config.add_control_adapter is True
assert config.cam_method == "prope"
assert config.attn_compress == 1
def test_dreamx_world_converted_5b_transformer_strict_loads():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
_load_fastvideo_transformer_strict(device, torch.bfloat16)
def test_dreamx_world_fastvideo_prope_branch_smoke():
block = DreamXWorldTransformerBlock(
8,
32,
1,
cross_attn_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
layer_idx=0,
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA,),
)
assert block.cam_self_attn is not None
assert [
name for name, _ in block.named_parameters()
if name.startswith("cam_self_attn.")
][:8] == [
"cam_self_attn.q_proj.weight",
"cam_self_attn.q_proj.bias",
"cam_self_attn.k_proj.weight",
"cam_self_attn.k_proj.bias",
"cam_self_attn.v_proj.weight",
"cam_self_attn.v_proj.bias",
"cam_self_attn.out_proj.weight",
"cam_self_attn.out_proj.bias",
]
module = DreamXPropeSelfAttention(
dim=8,
attn_dim=8,
num_heads=1,
qk_norm="rms_norm_across_heads",
).eval()
assert module.num_heads == 1
assert module.head_dim == 8
assert tuple(module.out_proj.weight.shape) == (8, 8)
assert torch.count_nonzero(module.out_proj.weight) == 0
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for transformer parity.")
def test_dreamx_world_transformer_parity_scaffold(monkeypatch):
device = torch.device("cuda:0")
dtype = torch.float32
inputs = _make_inputs(device, dtype)
official = _load_official_transformer(device, dtype)
official_out = _run_official(official, inputs)
del official
torch.cuda.empty_cache()
import fastvideo.attention.layer as attention_layer
import fastvideo.distributed.communication_op as communication_op
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
monkeypatch.setattr(attention_layer, "get_sp_parallel_rank", lambda: 0)
monkeypatch.setattr(attention_layer, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_to_all_4D", lambda tensor, scatter_dim=2, gather_dim=1: tensor)
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_gather", lambda tensor, dim=-1: tensor)
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(communication_op, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(fastvideo_dreamx, "sequence_model_parallel_shard", lambda tensor, dim=1: (tensor, tensor.shape[dim]))
monkeypatch.setattr(
fastvideo_dreamx,
"sequence_model_parallel_all_gather_with_unpad",
lambda tensor, original_seq_len, dim=1: tensor.narrow(dim, 0, original_seq_len),
)
fastvideo = _load_fastvideo_transformer(device, dtype)
fastvideo_out = _run_fastvideo(fastvideo, inputs)
assert official_out.shape == fastvideo_out.shape
diff = (official_out - fastvideo_out).abs()
print(f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
assert_close(fastvideo_out, official_out, atol=1e-1, rtol=1e-1)
@@ -0,0 +1,220 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World Wan2.2 VAE reuse parity scaffold.
Coverage scope: implementation_subcomponent. The official side uses
AutoencoderKLWan3_8 from DreamX, while the FastVideo side targets the native
Wan VAE. This remains a scaffold until Wan2.2 base VAE weights are staged.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import re
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_vae_config
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
PARITY_SCOPE = "implementation_subcomponent"
def _add_official_to_path():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _map_residual_subkey(prefix: str, sub: str) -> str | None:
if sub == "residual.0.gamma":
return f"{prefix}.norm1.gamma"
match = re.match(r"^residual\.2\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv1.{match.group(1)}"
if sub == "residual.3.gamma":
return f"{prefix}.norm2.gamma"
match = re.match(r"^residual\.6\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv2.{match.group(1)}"
match = re.match(r"^shortcut\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv_shortcut.{match.group(1)}"
return None
def _map_attention_subkey(prefix: str, sub: str) -> str | None:
if sub == "norm.gamma":
return f"{prefix}.norm.gamma"
match = re.match(r"^(to_qkv|proj)\.(weight|bias)$", sub)
if match:
return f"{prefix}.{match.group(1)}.{match.group(2)}"
return None
def _map_resample_subkey(prefix: str, sub: str) -> str | None:
match = re.match(r"^resample\.1\.(weight|bias)$", sub)
if match:
return f"{prefix}.resample.1.{match.group(1)}"
match = re.match(r"^time_conv\.(weight|bias)$", sub)
if match:
return f"{prefix}.time_conv.{match.group(1)}"
return None
def _map_dreamx_raw_vae_key(key: str) -> str | None:
match = re.match(r"^(conv1|conv2)\.(weight|bias)$", key)
if match:
prefix = "quant_conv" if match.group(1) == "conv1" else "post_quant_conv"
return f"{prefix}.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if match:
return f"{match.group(1)}.conv_in.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if match:
return f"{match.group(1)}.norm_out.gamma"
match = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if match:
return f"{match.group(1)}.conv_out.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if match:
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.0", match.group(2))
match = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if match:
return _map_attention_subkey(f"{match.group(1)}.mid_block.attentions.0", match.group(2))
match = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if match:
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.1", match.group(2))
match = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
if match:
stage = int(match.group(1))
block = int(match.group(2))
sub = match.group(3)
if block in (0, 1):
return _map_residual_subkey(f"encoder.down_blocks.{stage}.resnets.{block}", sub)
if block == 2:
return _map_resample_subkey(f"encoder.down_blocks.{stage}.downsampler", sub)
return None
match = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
if match:
stage = int(match.group(1))
block = int(match.group(2))
sub = match.group(3)
if block in (0, 1, 2):
return _map_residual_subkey(f"decoder.up_blocks.{stage}.resnets.{block}", sub)
if block == 3:
return _map_resample_subkey(f"decoder.up_blocks.{stage}.upsampler", sub)
return None
return None
def _vae_kwargs():
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
return OmegaConf.to_container(config["vae_kwargs"])
def _load_official_vae(device, dtype):
_add_official_to_path()
_install_xfuser_stub()
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
if not vae_path.exists():
pytest.skip(f"Wan2.2 base VAE weights missing: {vae_path}")
try:
from models import AutoencoderKLWan3_8
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX VAE: {exc}")
model = AutoencoderKLWan3_8.from_pretrained(str(vae_path), additional_kwargs=_vae_kwargs())
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_vae(device, dtype):
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
if not vae_path.exists():
pytest.skip(f"Wan2.2 raw VAE weights missing: {vae_path}")
config = make_dreamx_world_5b_cam_vae_config()
config.load_encoder = True
config.load_decoder = True
model = AutoencoderKLWan(config).to(device=device, dtype=dtype)
raw_state = torch.load(str(vae_path), map_location="cpu", weights_only=True)
mapped_state = {}
for key, value in raw_state.items():
mapped_key = _map_dreamx_raw_vae_key(key)
if mapped_key is None:
raise AssertionError(f"Unmapped DreamX raw VAE key: {key}")
mapped_state[mapped_key] = value
model.load_state_dict(mapped_state, strict=True)
return model.eval()
def _normalize_fastvideo_vae_latent(latent: torch.Tensor) -> torch.Tensor:
config = make_dreamx_world_5b_cam_vae_config()
mean = torch.tensor(config.latents_mean, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(config.latents_std, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
return (latent - mean) / std
def test_dreamx_world_vae_config_matches_wan22_shape():
config = make_dreamx_world_5b_cam_vae_config()
assert config.z_dim == 48
assert config.in_channels == 12
assert config.out_channels == 12
assert config.base_dim == 160
assert config.decoder_base_dim == 256
assert config.scale_factor_temporal == 4
assert config.scale_factor_spatial == 16
assert config.patch_size == 2
assert config.is_residual is True
assert config.clip_output is False
assert len(config.latents_mean) == 48
assert len(config.latents_std) == 48
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE parity.")
def test_dreamx_world_vae_encode_parity_scaffold():
device = torch.device("cuda:0")
dtype = torch.bfloat16
official = _load_official_vae(device, dtype)
fastvideo = _load_fastvideo_vae(device, dtype)
torch.manual_seed(123)
video = torch.randn(1, 3, 5, 64, 64, device=device, dtype=dtype).clamp(-1, 1)
with torch.inference_mode():
official_latent = official.encode(video).latent_dist.mean.float().cpu()
fastvideo_latent = _normalize_fastvideo_vae_latent(fastvideo.encode(video).mean).float().cpu()
assert official_latent.shape == fastvideo_latent.shape
assert_close(fastvideo_latent, official_latent, atol=5e-2, rtol=5e-2)
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline parity checks.
This test compares the FastVideo pipeline's DreamX-specific conditioning and
single-step scheduler path against an explicit hand-rolled pass using the same
loaded modules. It is intentionally local and deterministic: component parity
against the official DreamX repository lives in ``tests/local_tests/dreamx_world``.
"""
from __future__ import annotations
import gc
import os
from pathlib import Path
from typing import Any, cast
import pytest
import torch
from torch.testing import assert_close
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
# Local converted dir or HF repo id; the loader downloads hub ids itself.
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="DreamX-World pipeline parity requires CUDA",
)
def _run_worker_forward_batch(worker_wrapper: Any, request_kwargs: dict[str, Any]) -> torch.Tensor:
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import shallow_asdict
fastvideo_args = worker_wrapper.worker.fastvideo_args
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
sampling_param.update({
key: value
for key, value in request_kwargs.items()
if key not in {"prompt", "output_path"}
})
sampling_param.prompt = request_kwargs["prompt"]
latents_size = [
(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8,
]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
output_batch = worker_wrapper.worker.pipeline.forward(batch, fastvideo_args)
assert output_batch.output is not None
return output_batch.output.detach().cpu()
def _close_generator(generator: Any) -> None:
generator.shutdown()
gc.collect()
torch.cuda.empty_cache()
def test_dreamx_world_one_step_pipeline_latent_matches_manual_pass() -> None:
from fastvideo import VideoGenerator
common_kwargs = dict(
prompt="a quiet road through a futuristic city at sunrise",
output_path="outputs_video/dreamx_world_parity",
save_video=False,
return_frames=True,
height=64,
width=64,
num_frames=9,
num_inference_steps=1,
guidance_scale=1.0,
action_list=["w"],
action_speed_list=[2.0],
seed=123,
)
generator = VideoGenerator.from_pretrained(
MODEL_DIR,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
output_type="latent",
override_pipeline_cls_name="DreamXWorldPipeline",
)
try:
result = cast(dict[str, Any], generator.generate_video(**common_kwargs))
pipeline_latents = cast(torch.Tensor, result["samples"]).detach().cpu()
manual_latents = generator.executor.collective_rpc(
_run_worker_forward_batch,
kwargs={"request_kwargs": common_kwargs},
)[0]
finally:
_close_generator(generator)
assert_close(pipeline_latents, manual_latents, atol=0.0, rtol=0.0)
@@ -0,0 +1,157 @@
# SPDX-License-Identifier: Apache-2.0
"""Smoke tests for the DreamX-World-5B-Cam pipeline."""
from __future__ import annotations
import os
from pathlib import Path
from typing import Any, cast
import pytest
import torch
from PIL import Image, ImageDraw
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
# Local converted dir or HF repo id; the loader downloads hub ids itself.
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
def _write_smoke_image(path: Path) -> None:
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
draw = ImageDraw.Draw(image)
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
image.save(path)
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="DreamX-World pipeline smoke requires CUDA",
)
def test_dreamx_world_typed_surface_preflight() -> None:
import fastvideo.registry as registry
from fastvideo.api.presets import get_preset, get_presets_for_family
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import (
DreamXWorldPipeline,
EntryClass,
)
assert DreamXWorldPipeline.__name__ == "DreamXWorldPipeline"
assert EntryClass is DreamXWorldPipeline
assert DreamXWorldPipeline._required_config_modules == [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
default_preset, model_family = registry.get_preset_selection(
"GD-ML/DreamX-World-5B-Cam"
)
assert model_family == "dreamx_world"
assert default_preset == "dreamx_world_5b_cam"
info = registry.get_model_info(
"GD-ML/DreamX-World-5B-Cam",
workload_type=WorkloadType.I2V,
override_pipeline_cls_name="DreamXWorldPipeline",
)
assert info.pipeline_cls is DreamXWorldPipeline
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
names = {p.name for p in get_presets_for_family("dreamx_world")}
assert "dreamx_world_5b_cam" in names
preset = get_preset("dreamx_world_5b_cam", "dreamx_world")
assert preset.defaults["num_inference_steps"] == 30
assert preset.defaults["height"] == 480
assert preset.defaults["width"] == 832
assert preset.defaults["num_frames"] == 161
assert preset.defaults["guidance_scale"] == 5.0
cfg = DreamXWorld5BCamPipelineConfig()
assert cfg.flow_shift == 3.0
assert cfg.ti2v_task is True
assert cfg.expand_timesteps is True
assert cfg.dit_config.arch_config.add_control_adapter is True
assert cfg.dit_config.arch_config.cam_method == "prope"
def test_dreamx_world_camera_stage_writes_y_camera() -> None:
from types import SimpleNamespace
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
batch = ForwardBatch(
data_type="video",
prompt="camera smoke",
latents=torch.zeros(1, 48, 3, 8, 8, dtype=torch.bfloat16, device="cuda"),
num_frames=9,
height=64,
width=64,
action_list=["w", "d"],
action_speed_list=[2.0, 1.0],
)
out = DreamXWorldCameraConditioningStage().forward(batch, cast(Any, SimpleNamespace()))
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
assert set(y_camera) == {"viewmats", "K"}
assert y_camera["viewmats"].shape == (1, 3, 4, 4)
assert y_camera["K"].shape == (1, 3, 3, 3)
assert y_camera["viewmats"].device.type == "cuda"
assert y_camera["viewmats"].dtype == torch.bfloat16
def test_dreamx_world_pipeline_load_generate_latent_smoke(tmp_path: Path) -> None:
from fastvideo import VideoGenerator
image_path = tmp_path / "dreamx_world_smoke_input.png"
_write_smoke_image(image_path)
generator = VideoGenerator.from_pretrained(
MODEL_DIR,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
output_type="latent",
override_pipeline_cls_name="DreamXWorldPipeline",
)
try:
result = generator.generate_video(
prompt="a quiet road through a futuristic city at sunrise",
output_path="outputs_video/dreamx_world_smoke",
save_video=False,
return_frames=True,
height=64,
width=64,
num_frames=9,
num_inference_steps=1,
guidance_scale=1.0,
image_path=str(image_path),
action_list=["w"],
action_speed_list=[2.0],
seed=0,
)
finally:
generator.shutdown()
assert isinstance(result, dict)
samples = cast(dict[str, Any], result)["samples"]
assert torch.is_tensor(samples)
assert samples.ndim == 5
assert samples.shape[1] == 48
assert torch.isfinite(samples).all()