Compare commits
21
Commits
v2
...
maint/pr1538-wave5
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0e727f5ee | ||
|
|
de6a60ac90 | ||
|
|
1b4fdc2c41 | ||
|
|
385cd65abb | ||
|
|
156d86b70c | ||
|
|
bab79fb5f6 | ||
|
|
073a9e78f2 | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -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",
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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" && \
|
||||
|
||||
@@ -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` |
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -74,6 +74,23 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -0,0 +1,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()
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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"),
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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 {
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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))
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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. |
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user