Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
30e45c2411 | ||
|
|
2a4fe697a6 | ||
|
|
921db7479d | ||
|
|
7f539424cb | ||
|
|
19a838f54f | ||
|
|
d922ab2cbc | ||
|
|
9ea77d37f3 | ||
|
|
2e35b0c6bd | ||
|
|
1c627a3f98 | ||
|
|
a931efe33a | ||
|
|
041e5e9029 | ||
|
|
efcc245c2e | ||
|
|
922e7e0813 | ||
|
|
c62a8514b0 |
@@ -8,4 +8,3 @@
|
||||
{"name": "decompose-pipeline-pr", "description": "Decompose an oversized FastVideo pipeline PR into a stack of independently-reviewable PRs. Tiers the diff by blast radius (invisible / dead code / cross-cutting infra / activation), produces a branch graph and worktree bootstrap, drafts the AGENTS.md manifest, flags missing tests on cross-cutting infra changes, and extracts lessons from the PR body. Worked example: PR #1280 daVinci-MagiHuman (9.8k LOC) decomposed into 10 stacked PRs.", "path": "decompose-pipeline-pr/SKILL.md", "status": "tested", "trust": "medium"}
|
||||
{"name": "reseed-performance-baseline", "description": "Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline must be advanced by replicating one reviewed shifted source result into three success=true records, or five records when explicitly requested", "path": "reseed-performance-baseline/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "release", "description": "Cut a new FastVideo release. Bumps the version across the three authoritative files (fastvideo/version.py, pyproject.toml, pyproject_other.toml), opens a [chore]: release PR, and documents the post-merge tag + GitHub release ritual. Triggers on requests like \"release X.Y.Z\", \"cut a release\", \"bump version to X.Y.Z\", \"publish to PyPI\".", "path": "release/SKILL.md", "status": "draft", "trust": "low"}
|
||||
|
||||
@@ -1,248 +0,0 @@
|
||||
---
|
||||
name: release
|
||||
description: Cut a new FastVideo release. Bumps the version across the three authoritative files (fastvideo/version.py, pyproject.toml, pyproject_other.toml), opens a [chore]: release PR, and documents the post-merge tag + GitHub release ritual. Triggers on requests like "release X.Y.Z", "cut a release", "bump version to X.Y.Z", "publish to PyPI".
|
||||
---
|
||||
|
||||
# FastVideo release skill
|
||||
|
||||
End-to-end recipe for cutting a FastVideo release. The PyPI publish is automatic — pushing a `pyproject.toml` version change to `main` triggers `.github/workflows/publish-fastvideo.yml`. Your job is to land the version bump cleanly and follow up with a git tag + GitHub Release for the changelog.
|
||||
|
||||
## Inputs
|
||||
|
||||
- `${NEW}` — the new version (e.g. `0.2.0`). Required.
|
||||
- `${OLD}` — the current version. Auto-detect with: `grep -oP '__version__ = "\K[^"]+' fastvideo/version.py` from the repo root.
|
||||
|
||||
## When to use
|
||||
|
||||
Trigger phrases: "release X.Y.Z", "cut a release", "bump version to X.Y.Z", "publish to PyPI", "tag a release".
|
||||
|
||||
## Files to update (3 — the authoritative list)
|
||||
|
||||
These are the ONLY files that carry the version as a Python/package declaration:
|
||||
|
||||
| File | Line | Change |
|
||||
|---|---|---|
|
||||
| `fastvideo/version.py` | 1 | `__version__ = "${OLD}"` → `__version__ = "${NEW}"` |
|
||||
| `pyproject.toml` | 7 | `version = "${OLD}"` → `version = "${NEW}"` |
|
||||
| `pyproject_other.toml` | 7 | `version = "${OLD}"` → `version = "${NEW}"` |
|
||||
|
||||
`fastvideo/__init__.py` re-exports `__version__` from `fastvideo.version`, so no edit needed there.
|
||||
|
||||
## Files NOT to touch
|
||||
|
||||
- `apps/dreamverse/pyproject.toml` — declares `"fastvideo>=X.Y.Z"` as a floor. A new release usually still satisfies the floor; bumping it is a separate policy call (does dreamverse strictly require the new version?). Leave alone unless explicitly asked.
|
||||
- `.agents/memory/**/*.md` — historical notes; the version strings in there are snapshots, not declarations.
|
||||
- `examples/`, `docs/` — version mentions are illustrative; not authoritative.
|
||||
- `uv.lock` — main does NOT track a `uv.lock`. Do NOT run `uv lock` as part of a release.
|
||||
|
||||
## Workflow
|
||||
|
||||
### 1. Verify clean state
|
||||
|
||||
```bash
|
||||
# From the primary FastVideo jj workspace
|
||||
jj git fetch
|
||||
OLD=$(grep -oP '__version__ = "\K[^"]+' fastvideo/version.py)
|
||||
echo "current: $OLD → target: $NEW"
|
||||
```
|
||||
|
||||
Confirm `$NEW > $OLD` follows semver. Check prior tags for the pattern:
|
||||
```bash
|
||||
gh release list --repo hao-ai-lab/FastVideo --limit 5
|
||||
```
|
||||
|
||||
### 2. Create a dedicated jj workspace + bookmark
|
||||
|
||||
```bash
|
||||
WS=/home/william5lin/FastVideo_release_${NEW//./_}
|
||||
jj workspace add --name release-${NEW//./-} "$WS"
|
||||
cd "$WS"
|
||||
jj new main@origin -m "[chore]: release v${NEW}"
|
||||
jj bookmark create chore/release-${NEW} -r @
|
||||
```
|
||||
|
||||
### 3. Apply the 3-file bump
|
||||
|
||||
Use the `edit` tool or `sed -i` with exact context. Example with sed:
|
||||
```bash
|
||||
sed -i "s/__version__ = \"${OLD}\"/__version__ = \"${NEW}\"/" fastvideo/version.py
|
||||
sed -i "0,/version = \"${OLD}\"/s//version = \"${NEW}\"/" pyproject.toml
|
||||
sed -i "0,/version = \"${OLD}\"/s//version = \"${NEW}\"/" pyproject_other.toml
|
||||
```
|
||||
(The `0,/.../s//.../` form replaces only the FIRST match in each `pyproject*.toml`, since `${OLD}` might appear elsewhere as a constraint.)
|
||||
|
||||
### 4. Verify
|
||||
|
||||
```bash
|
||||
jj diff --name-only -r @ # MUST be exactly 3 files
|
||||
jj diff --stat -r @ # MUST be +3 / -3
|
||||
grep -nE "${OLD//./\\.}" fastvideo/version.py pyproject.toml pyproject_other.toml
|
||||
# expect NO matches in the three files
|
||||
```
|
||||
|
||||
### 5. Lint
|
||||
|
||||
```bash
|
||||
pre-commit run --files fastvideo/version.py pyproject.toml pyproject_other.toml
|
||||
```
|
||||
Must pass. Never `--no-verify`.
|
||||
|
||||
### 6. Describe + push
|
||||
|
||||
```bash
|
||||
jj describe -m "[chore]: release v${NEW}
|
||||
|
||||
Bumps FastVideo from ${OLD} to ${NEW}.
|
||||
|
||||
Files updated:
|
||||
fastvideo/version.py
|
||||
pyproject.toml
|
||||
pyproject_other.toml
|
||||
|
||||
Note: pushing this to main triggers .github/workflows/publish-fastvideo.yml,
|
||||
which detects the pyproject.toml version change and publishes to PyPI.
|
||||
Tag v${NEW} + GitHub release notes follow merge."
|
||||
|
||||
jj git push --bookmark chore/release-${NEW}
|
||||
```
|
||||
|
||||
### 7. Open PR
|
||||
|
||||
```bash
|
||||
gh pr create \
|
||||
--repo hao-ai-lab/FastVideo \
|
||||
--base main \
|
||||
--head chore/release-${NEW} \
|
||||
--title "[chore]: release v${NEW}" \
|
||||
--body-file - <<EOF
|
||||
## Summary
|
||||
|
||||
Bumps FastVideo from \`${OLD}\` to \`${NEW}\`.
|
||||
|
||||
## Files updated (3)
|
||||
|
||||
- \`fastvideo/version.py\`
|
||||
- \`pyproject.toml\`
|
||||
- \`pyproject_other.toml\`
|
||||
|
||||
## Out of scope
|
||||
|
||||
\`apps/dreamverse/pyproject.toml\` floor (\`fastvideo>=${OLD}\`) — \`${NEW}\` satisfies it; bumping is a separate policy call.
|
||||
|
||||
## After merge
|
||||
|
||||
\`.github/workflows/publish-fastvideo.yml\` auto-publishes to PyPI on push-to-main when \`pyproject.toml\` changes.
|
||||
|
||||
Manual follow-up:
|
||||
- Tag the merge commit: \`git tag v${NEW} <merge-sha> && git push origin v${NEW}\`
|
||||
- Create GitHub Release \`v${NEW}\` matching the prior \`Release X.Y.Z\` pattern.
|
||||
EOF
|
||||
```
|
||||
|
||||
### 8. Post-merge ritual (do AFTER the PR merges)
|
||||
|
||||
1. **Tag the merge commit**:
|
||||
```bash
|
||||
git fetch origin
|
||||
MERGE_SHA=$(gh pr view <PR-NUMBER> --repo hao-ai-lab/FastVideo --json mergeCommit --jq .mergeCommit.oid)
|
||||
git tag v${NEW} ${MERGE_SHA}
|
||||
git push origin v${NEW}
|
||||
```
|
||||
2. **Confirm PyPI publish workflow ran**:
|
||||
```bash
|
||||
gh run list --repo hao-ai-lab/FastVideo --workflow publish-fastvideo.yml --limit 3
|
||||
```
|
||||
3. **Create the GitHub Release**:
|
||||
```bash
|
||||
gh release create v${NEW} \
|
||||
--repo hao-ai-lab/FastVideo \
|
||||
--title "Release ${NEW}" \
|
||||
--notes "<changelog highlights — what shipped since v${OLD}>" \
|
||||
--target main
|
||||
```
|
||||
Use `gh release view v${OLD}` to mirror tone/structure from the prior release.
|
||||
4. **Cleanup**: after merge + tag + release land, tear down the workspace:
|
||||
```bash
|
||||
jj workspace forget release-${NEW//./-}
|
||||
rm -rf "$WS"
|
||||
jj bookmark delete chore/release-${NEW}
|
||||
```
|
||||
|
||||
## Verification gates (must all pass before pushing)
|
||||
|
||||
- `jj diff --name-only -r @` returns exactly 3 files
|
||||
- `jj diff --stat -r @` shows `+3 / -3`
|
||||
- `grep -E "${OLD//./\\.}" fastvideo/version.py pyproject.toml pyproject_other.toml` returns no matches
|
||||
- `pre-commit run --files <the-three>` passes
|
||||
- No `uv.lock` in the change
|
||||
- No source-code files touched
|
||||
|
||||
## Conventions (enforced)
|
||||
|
||||
- Commit subject: `[chore]: release v${NEW}` (under 72 chars).
|
||||
- NEVER add AI co-author trailers (`Co-Authored-By: Claude`, "Generated with…", etc.).
|
||||
- NEVER `--no-verify`.
|
||||
- NEVER `uv lock` as part of a release — main doesn't track the lockfile.
|
||||
- Tag format: `vX.Y.Z` (with leading `v`), matching prior releases.
|
||||
|
||||
## Why three files?
|
||||
|
||||
`pyproject.toml` and `pyproject_other.toml` are two co-existing project metadata files (the project ships both — the latter is a slimmer variant without dreamverse/job-runner extras). Both carry an authoritative `version = "X.Y.Z"` field and must stay in lock-step. `fastvideo/version.py` is the runtime source of truth re-exported by `fastvideo/__init__.py`.
|
||||
|
||||
## Publish workflow contract
|
||||
|
||||
`.github/workflows/publish-fastvideo.yml` triggers on `push` to `main` when `pyproject.toml` changes. It compares the new `version` field to the previous commit's `version` field and, if different, builds + publishes to PyPI. The version bump in `pyproject_other.toml` does NOT trigger the workflow (only `pyproject.toml` is in the `paths:` filter), but keeping the two in sync prevents installer surprises for users of the alternate metadata file.
|
||||
|
||||
## PyPI publish failure modes
|
||||
|
||||
The publish workflow ran on the merge commit but the PyPI upload can still fail at the OIDC trusted-publishing exchange. Always verify the workflow succeeded — do not assume "merge implies published":
|
||||
|
||||
```bash
|
||||
gh run list --repo hao-ai-lab/FastVideo --workflow publish-fastvideo.yml --limit 3
|
||||
```
|
||||
|
||||
Look for the run on the release merge commit. If it shows `failure`, dump the failed log:
|
||||
|
||||
```bash
|
||||
gh run view <run-id> --repo hao-ai-lab/FastVideo --log-failed | tail -80
|
||||
```
|
||||
|
||||
### Known failure: `invalid-publisher` (Trusted Publisher claim mismatch)
|
||||
|
||||
The most common failure surfaces as:
|
||||
|
||||
```
|
||||
Trusted publishing exchange failure:
|
||||
* `invalid-publisher`: valid token, but no corresponding publisher
|
||||
(Publisher with matching claims was not found)
|
||||
* environment: MISSING
|
||||
```
|
||||
|
||||
This means the PyPI Trusted Publisher registered for the project expects an `environment` claim (e.g. `pypi`) that the workflow job does not set. Two recovery paths:
|
||||
|
||||
**A. Fix the trusted publisher + re-run the workflow** (cleaner long-term):
|
||||
1. On `pypi.org/manage/project/fastvideo/settings/publishing/`, either remove the `Environment name` field from the registered publisher, OR add `environment: pypi` (matching the existing PyPI config) to the `build-publish-main` job in `.github/workflows/publish-fastvideo.yml`.
|
||||
2. Re-run the failed workflow:
|
||||
```bash
|
||||
gh run rerun <run-id> --repo hao-ai-lab/FastVideo --failed
|
||||
```
|
||||
|
||||
**B. Manual one-shot publish** (faster, no infra change):
|
||||
|
||||
```bash
|
||||
git checkout <merge-sha> # the v${NEW} merge commit on main
|
||||
uv build # builds sdist + wheel into dist/
|
||||
uv publish --token <PYPI_TOKEN> # or: twine upload dist/*
|
||||
```
|
||||
|
||||
PyPI is **immutable per version** — if any artifact for `${NEW}` got uploaded (sdist or wheel), you cannot re-upload it. Check before retrying:
|
||||
|
||||
```bash
|
||||
curl -s https://pypi.org/pypi/fastvideo/${NEW}/json | python3 -c "import sys,json; d=json.load(sys.stdin); print('on pypi:', list(d['urls'][0].keys()) if d.get('urls') else 'NOT_PUBLISHED')"
|
||||
```
|
||||
|
||||
If `NOT_PUBLISHED`, either recovery path works. If anything is already up, you have to cut a `${NEW}.postN` patch release instead.
|
||||
|
||||
### Tag and GitHub Release are independent
|
||||
|
||||
The `git tag v${NEW}` and `gh release create v${NEW}` steps are **independent of PyPI publish success**. If you created the tag + release before noticing the publish failure, that's fine — keep them; just complete the PyPI publish via path A or B above. Do NOT delete and re-create the tag, because doing so will cause confusion in dependents that pin to the tag.
|
||||
@@ -34,6 +34,8 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
|
||||
@@ -55,12 +55,14 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install FastVideo Unified Kernel.
|
||||
# This build machine has no GPU, so target Hopper (sm_90a) explicitly instead
|
||||
# of probing a live device for the arch (matches the released kernel wheel).
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
|
||||
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -55,12 +55,14 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install FastVideo Unified Kernel.
|
||||
# This build machine has no GPU, so target Hopper (sm_90a) explicitly instead
|
||||
# of probing a live device for the arch (matches the released kernel wheel).
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
|
||||
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -55,11 +55,13 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install FastVideo Unified Kernel.
|
||||
# This build machine has no GPU, so target Hopper (sm_90a) explicitly instead
|
||||
# of probing a live device for the arch (matches the released kernel wheel).
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -55,12 +55,14 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install FastVideo Unified Kernel.
|
||||
# This build machine has no GPU, so target Hopper (sm_90a) explicitly instead
|
||||
# of probing a live device for the arch (matches the released kernel wheel).
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
|
||||
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -108,6 +108,7 @@ surfaces:
|
||||
vae_sp: generator.pipeline.preset_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.preset_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.preset_overrides.ti2v_task
|
||||
lucy_edit_task: generator.pipeline.preset_overrides.lucy_edit_task
|
||||
boundary_ratio: generator.pipeline.preset_overrides.boundary_ratio
|
||||
compatibility_only:
|
||||
model_path: "Redundant with generator.model_path."
|
||||
@@ -125,9 +126,18 @@ surfaces:
|
||||
text_encoder_configs: "Legacy internal component config object."
|
||||
preprocess_text_funcs: "Internal text preprocessing hooks."
|
||||
postprocess_text_funcs: "Internal text postprocessing hooks."
|
||||
scheduler_step_in_fp32: "Runtime scheduler precision toggle; not part of the public typed inference API."
|
||||
|
||||
pipeline_config_extensions:
|
||||
preset_owned:
|
||||
flux2_text_encoder_type:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.flux_2.Flux2PipelineConfig
|
||||
- fastvideo.configs.pipelines.flux_2.Flux2KleinPipelineConfig
|
||||
text_encoder_out_layers:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.flux_2.Flux2PipelineConfig
|
||||
- fastvideo.configs.pipelines.flux_2.Flux2KleinPipelineConfig
|
||||
conditioning_strategy:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
@@ -491,6 +501,8 @@ surfaces:
|
||||
inpaint_mask: request.extensions.stable_audio.inpaint_mask
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
latents: "Pre-generated diffusion latents supplied by parity/debug harnesses; not a public input."
|
||||
max_sequence_length: "Model-specific text-encoder sequence cap; not part of the public typed inference API."
|
||||
|
||||
sampling_param_extensions: {}
|
||||
|
||||
|
||||
@@ -58,6 +58,7 @@ 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 | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| 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 | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
|
||||
@@ -78,6 +79,9 @@ pipeline initialization and sampling.
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
|
||||
focused on inference integration for video editing workflows.
|
||||
|
||||
`Sliding Tile Attn (Legacy Branch)` entries refer to the archived
|
||||
`sta_do_not_delete` branch workflow, not active `main` inference wiring.
|
||||
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run full Flux2 text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have a local or HF Diffusers-format full Flux2 checkpoint and want a
|
||||
minimal text-to-image generation command that uses embedded guidance."
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
ComponentConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run full Flux2 text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="black-forest-labs/FLUX.2-dev",
|
||||
help="HF id or local diffusers-format full Flux2 weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="outputs/flux2/flux2.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="a photo of a banana on a wooden table, studio lighting",
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=4.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=None)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND, for example TORCH_SDPA.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (
|
||||
args.num_gpus if args.num_gpus > 1 else 1
|
||||
)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (
|
||||
1 if args.num_gpus > 1 else args.num_gpus
|
||||
)
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
workload_type="t2i",
|
||||
components=ComponentConfig(override_pipeline_cls_name="Flux2Pipeline"),
|
||||
),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
sampling = SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
)
|
||||
extensions = {}
|
||||
if args.max_sequence_length is not None:
|
||||
extensions["max_sequence_length"] = args.max_sequence_length
|
||||
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=sampling,
|
||||
output=OutputConfig(
|
||||
output_path=str(output),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
extensions=extensions,
|
||||
)
|
||||
generator.generate(request)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,98 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run Flux2 Klein text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I need a short local smoke for the Flux2 Klein checkpoint before wiring it
|
||||
into an image workflow. Use the model's distilled four-step defaults and
|
||||
write a single PNG so I can compare the output against the reference."
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run Flux2 Klein text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="black-forest-labs/FLUX.2-klein-4B",
|
||||
help="HF id or local diffusers-format Flux2 Klein weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
default="outputs/flux2/flux2_klein.png",
|
||||
help="PNG output path or output directory.",
|
||||
)
|
||||
parser.add_argument("--prompt", default=DEFAULT_PROMPT, help="Prompt text.")
|
||||
parser.add_argument("--seed", type=int, default=0, help="Generation seed.")
|
||||
parser.add_argument("--height", type=int, default=1024, help="Output image height.")
|
||||
parser.add_argument("--width", type=int, default=1024, help="Output image width.")
|
||||
parser.add_argument("--steps", type=int, default=4, help="Number of denoising steps.")
|
||||
parser.add_argument("--num-gpus", type=int, default=1, help="Number of GPUs to use.")
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND, for example TORCH_SDPA.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
vae=True,
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=1.0,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
),
|
||||
)
|
||||
generator.generate(request)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,289 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video with torch.compile + timing breakdown.
|
||||
|
||||
This example runs the LTX-2.3 distilled student model on a single GPU with
|
||||
torch.compile fully enabled, then prints a per-stage timing breakdown so the
|
||||
user can see where wall-time goes. It is meant as the canonical entry point
|
||||
for trying out the LTX-2.3 i2v path on `hao-ai-lab/FastVideo:main`.
|
||||
|
||||
Quick start
|
||||
-----------
|
||||
export LTX23_I2V_IMAGE=/path/to/your/portrait_or_product.jpg
|
||||
# optional overrides:
|
||||
# export LTX23_I2V_PROMPT="a fashion model walks toward camera..."
|
||||
# export LTX23_OUTPUT_DIR=outputs_video/ltx2_3_distilled_i2v
|
||||
python examples/inference/basic/basic_ltx2_3_distilled_i2v.py
|
||||
|
||||
What the script does
|
||||
--------------------
|
||||
1. Loads FastVideo/LTX-2.3-Distilled-Diffusers (8 denoise + 3 refine steps,
|
||||
CFG=1, no refine LoRA — the distilled production recipe).
|
||||
2. Compiles the DiT, text encoder, and VAE (fullgraph, Inductor default
|
||||
mode — autotune adds ~7 min cold-compile here with no measurable
|
||||
e2e gain).
|
||||
3. Runs 2 warmup calls (untimed) + 2 measured calls. Two warmups are kept
|
||||
as a safety net — the first call pays cold compile + first-shape guard
|
||||
work, and a second warmup ensures any residual recompiles settle before
|
||||
we measure.
|
||||
4. Prints a per-stage breakdown and an average over the measured runs.
|
||||
|
||||
Hardware notes
|
||||
--------------
|
||||
- Single-GPU example; for multi-GPU sequence-parallel see the gradio demo
|
||||
under `examples/inference/gradio/local/gradio_local_demo_ltx2_3/`.
|
||||
- First-time compile takes ~30-40 min on GB200 (~20 min on H100; cached
|
||||
in `$TORCHINDUCTOR_CACHE_DIR` afterwards). Subsequent invocations only
|
||||
pay the one-time process load + a few seconds of dynamo trace.
|
||||
- On GB200 / Blackwell, run with `env -u LD_LIBRARY_PATH ...` to avoid a
|
||||
system-cuBLAS / torch-cuBLAS mismatch that fails every GEMM. The
|
||||
`_inductor.shape_padding = False` line below also avoids a pad_mm
|
||||
landmine on the same generation of cards.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch._inductor.config as _inductor
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
# Env knobs (set BEFORE importing fastvideo where possible — but
|
||||
# FASTVIDEO_ATTENTION_BACKEND is fine here because the worker reads it
|
||||
# on generator construction).
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
# Inductor knobs. The first one (shape_padding=False) is mandatory on
|
||||
# Blackwell to avoid a cuBLAS INVALID_VALUE crash inside pad_mm during
|
||||
# the refine path. The rest are autotune-friendliness flags.
|
||||
_inductor.shape_padding = False
|
||||
_inductor.conv_1x1_as_mm = True
|
||||
_inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v")
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
# Per-stage timing helpers --------------------------------------------------
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
"""Print stage execution times and return the sum, or None if missing."""
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
"""LTX-2.3 distilled snapshots ship a `spatial_upscaler/` subdir."""
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
# Main ---------------------------------------------------------------------
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
|
||||
# mode="default" — Inductor's default schedule matches max-autotune on
|
||||
# this pipeline (denoise/refine/decode all within ~5 ms, n=2) while
|
||||
# saving ~7 min of cold compile on a single GB200.
|
||||
torch_compile_kwargs = {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "default",
|
||||
"dynamic": False,
|
||||
}
|
||||
|
||||
# Loading the pipeline config *with model_path* binds model-specific
|
||||
# tuning (notably VAE precision/decoder defaults) into the config. Without
|
||||
# this, the generic pipeline config gives a substantially slower VAE
|
||||
# decode stage. `basic_ltx2_distilled_fast_profile.py` uses the same
|
||||
# pattern.
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = None
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
# LTX-2.3 distilled uses the two-stage refine pipeline; the refine
|
||||
# LoRA is intentionally empty for the distilled student.
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
# Compile the VAE codec submodules (encoder / decoder) too. The
|
||||
# `LTX2CausalVideoAutoencoder` declares `_compile_conditions` so
|
||||
# `_compile_with_conditions` targets just those submodules and
|
||||
# leaves the surrounding tiling control flow eager — needed for
|
||||
# fullgraph + dynamic=False to succeed. VAE eager decode is
|
||||
# ~1.0s; compiling it brings the stage to ~0.3s.
|
||||
enable_torch_compile_vae=True,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=torch_compile_kwargs,
|
||||
# Keep everything resident — no CPU offload for serving-style runs.
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
# i2v: anchor the input image at frame 0 with full strength.
|
||||
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
|
||||
# JPEG conditioning image.
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
warmup_runs = 2
|
||||
measured_runs = 2
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
# Warmup: untimed (but we still wall-clock them so the first compile
|
||||
# cost is visible to the reader).
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
# Cleanup warmup artifacts so the user only sees measured outputs.
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
# Measured.
|
||||
for m in range(measured_runs):
|
||||
out_path = OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_run_{m + 1}.mp4"
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
# Summary.
|
||||
print("\n=== summary ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,350 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video — typed API (``from_config`` / ``generate``).
|
||||
|
||||
Identical generation behavior to ``basic_ltx2_3_distilled_i2v.py``, but
|
||||
expressed through the newer typed surface (``GeneratorConfig`` /
|
||||
``GenerationRequest``) instead of the ``from_pretrained(**legacy_kwargs)``
|
||||
bridge. The typed API is now the preferred entry point — the legacy
|
||||
example still works but emits a ``DeprecationWarning`` for the LTX-2.3
|
||||
specific knobs.
|
||||
|
||||
Quick start
|
||||
-----------
|
||||
export LTX23_I2V_IMAGE=/path/to/your/portrait_or_product.jpg
|
||||
# optional overrides:
|
||||
# export LTX23_I2V_PROMPT="a fashion model walks toward camera..."
|
||||
# export LTX23_OUTPUT_DIR=outputs_video/ltx2_3_distilled_i2v_typed
|
||||
python examples/inference/basic/basic_ltx2_3_distilled_i2v_typed.py
|
||||
|
||||
What the script does
|
||||
--------------------
|
||||
1. Loads FastVideo/LTX-2.3-Distilled-Diffusers (8 denoise + 3 refine
|
||||
steps, CFG=1, no refine LoRA — the distilled production recipe).
|
||||
2. Compiles the DiT, text encoder, and VAE (fullgraph, Inductor default
|
||||
mode — autotune adds ~7 min cold-compile here with no measurable
|
||||
e2e gain).
|
||||
3. Runs 2 warmup calls (untimed) + 2 measured calls. Two warmups are
|
||||
kept as a safety net — the first call pays cold compile + first-shape
|
||||
guard work, and a second warmup ensures any residual recompiles
|
||||
settle before we measure.
|
||||
4. Prints a per-stage breakdown and an average over the measured runs.
|
||||
|
||||
Hardware notes
|
||||
--------------
|
||||
- Single-GPU example; for multi-GPU sequence-parallel see the gradio
|
||||
demo under ``examples/inference/gradio/local/gradio_local_demo_ltx2_3/``.
|
||||
- First-time compile takes ~30-40 min on GB200 (~20 min on H100;
|
||||
cached in ``$TORCHINDUCTOR_CACHE_DIR`` afterwards). Subsequent
|
||||
invocations only pay the one-time process load + a few seconds of
|
||||
dynamo trace.
|
||||
- On GB200 / Blackwell, run with ``env -u LD_LIBRARY_PATH ...`` to
|
||||
avoid a system-cuBLAS / torch-cuBLAS mismatch that fails every GEMM.
|
||||
The ``_inductor.shape_padding = False`` line below also avoids a
|
||||
``pad_mm`` landmine on the same generation of cards.
|
||||
|
||||
Typed-API mapping (legacy kwarg ↔ typed field)
|
||||
----------------------------------------------
|
||||
- ``num_gpus`` ↔ ``engine.num_gpus``
|
||||
- ``enable_torch_compile`` ↔ ``engine.compile.enabled``
|
||||
- ``enable_torch_compile_text_encoder`` ↔ ``engine.compile.text_encoder_enabled``
|
||||
- ``enable_torch_compile_vae`` ↔ ``engine.compile.vae_enabled``
|
||||
- ``torch_compile_kwargs`` ↔ ``engine.compile.backend/fullgraph/mode/dynamic``
|
||||
- ``torch_compile_kwargs_vae`` ↔ empty ``compile.vae_kwargs`` (inherits master)
|
||||
- ``dit_cpu_offload`` ↔ ``engine.offload.dit``
|
||||
- ``text_encoder_cpu_offload`` ↔ ``engine.offload.text_encoder``
|
||||
- ``vae_cpu_offload`` ↔ ``engine.offload.vae``
|
||||
- ``ltx2_vae_tiling`` ↔ ``pipeline.vae_tiling``
|
||||
- ``ltx2_refine_enabled`` ↔ ``pipeline.preset_overrides["refine"]["enabled"]``
|
||||
- ``ltx2_refine_upsampler_path`` ↔ ``pipeline.components.upsampler_weights``
|
||||
- ``ltx2_refine_lora_path`` ↔ ``pipeline.components.lora_path``
|
||||
- ``ltx2_refine_num_inference_steps`` ↔ ``pipeline.preset_overrides["refine"]["num_inference_steps"]``
|
||||
- ``ltx2_refine_guidance_scale`` ↔ ``pipeline.preset_overrides["refine"]["guidance_scale"]``
|
||||
- ``ltx2_refine_add_noise`` ↔ ``pipeline.preset_overrides["refine"]["add_noise"]``
|
||||
- ``pipeline_config=PipelineConfig.from_pretrained(model_root)`` ↔ (no-op — ``PipelineConfig.from_kwargs`` already resolves the model-specific class from ``model_path``)
|
||||
- ``pipeline_config.dit_config.quant_config = None`` ↔ leave ``engine.quantization`` unset
|
||||
- ``ltx2_images`` / ``ltx2_image_crf`` ↔ ``request.extensions`` (LTX-2 specific, no
|
||||
first-class typed field yet)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch._inductor.config as _inductor
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
ComponentConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
# Inductor knobs. ``shape_padding=False`` is mandatory on Blackwell to
|
||||
# avoid a cuBLAS INVALID_VALUE crash inside pad_mm during the refine
|
||||
# path. The rest are autotune-friendliness flags.
|
||||
_inductor.shape_padding = False
|
||||
_inductor.conv_1x1_as_mm = True
|
||||
_inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"
|
||||
)
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result, label: str) -> float | None:
|
||||
logging_info = getattr(result, "logging_info", None)
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = getattr(result, "logging_info", None)
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_typed.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
|
||||
# mode="default" — Inductor's default schedule matches max-autotune on
|
||||
# this pipeline (denoise/refine/decode all within ~5 ms, n=2) while
|
||||
# saving ~7 min of cold compile on a single GB200.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_root,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
# Keep DiT / text encoder / VAE resident on GPU — no CPU offload
|
||||
# for serving-style runs. ``image_encoder`` and
|
||||
# ``pin_cpu_memory`` are left at their schema defaults
|
||||
# (matches the legacy example, which only set these three).
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
text_encoder=False,
|
||||
vae=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=True,
|
||||
text_encoder_enabled=True,
|
||||
# ``vae_enabled`` triggers ``_compile_with_conditions`` on
|
||||
# ``LTX2CausalVideoAutoencoder``, which compiles just the
|
||||
# encoder/decoder submodules and leaves the surrounding
|
||||
# tiling control flow eager (required for ``fullgraph``).
|
||||
# Empty ``vae_kwargs`` → inherits the master kwargs below.
|
||||
vae_enabled=True,
|
||||
backend="inductor",
|
||||
fullgraph=True,
|
||||
mode="default",
|
||||
dynamic=False,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
# ``PipelineConfig.from_kwargs`` resolves the model-specific
|
||||
# pipeline-config class from ``model_path`` automatically, so we
|
||||
# don't need to set ``components.pipeline_config_path`` — the
|
||||
# model-specific VAE precision / decoder defaults are picked up
|
||||
# the same way the legacy example's
|
||||
# ``PipelineConfig.from_pretrained(model_root)`` did them.
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
# Distilled has no refine LoRA — omit ``lora_path``.
|
||||
),
|
||||
vae_tiling=False,
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
"enabled": True,
|
||||
"num_inference_steps": 3,
|
||||
"guidance_scale": 1.0,
|
||||
"add_noise": True,
|
||||
},
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
|
||||
def build_request(out_path: Path, seed: int) -> GenerationRequest:
|
||||
return GenerationRequest(
|
||||
prompt=PROMPT,
|
||||
# distilled is CFG-free; no negative prompt
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
num_videos_per_prompt=1,
|
||||
seed=seed,
|
||||
height=1280,
|
||||
width=832,
|
||||
num_frames=121,
|
||||
fps=24,
|
||||
num_inference_steps=8,
|
||||
guidance_scale=1.0,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(out_path),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
# LTX-2.3 i2v fields don't have first-class typed slots yet;
|
||||
# extensions is the documented bridge. ``ltx2_image_crf=0.0``
|
||||
# skips an extra JPEG re-encode of an already JPEG image.
|
||||
extensions={
|
||||
"ltx2_images": [(I2V_IMAGE, 0, 1.0)],
|
||||
"ltx2_image_crf": 0.0,
|
||||
},
|
||||
)
|
||||
|
||||
warmup_runs = 2
|
||||
measured_runs = 2
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
t0 = time.perf_counter()
|
||||
generator.generate(
|
||||
build_request(
|
||||
OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7
|
||||
)
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4"
|
||||
)
|
||||
print(
|
||||
f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}"
|
||||
)
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate(
|
||||
build_request(out_path, seed=2002 + m)
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
# ``e2e_latency`` is currently surfaced via ``result.extra``;
|
||||
# ``GenerationResult`` exposes ``generation_time`` as a
|
||||
# first-class field but the LTX-2 pipeline only fills the
|
||||
# legacy ``e2e_latency`` key. Prefer the explicit one, fall
|
||||
# back to wall-clock.
|
||||
e2e = (
|
||||
result.extra.get("e2e_latency")
|
||||
if hasattr(result, "extra") else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(
|
||||
f"[measured {m + 1}/{measured_runs}] "
|
||||
f"e2e={e2e:.2f}s wall={wall:.2f}s"
|
||||
)
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary ===")
|
||||
print(
|
||||
f"warmup wall-times: "
|
||||
f"{[round(x, 1) for x in warmup_secs]}"
|
||||
)
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,38 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_lucy_edit"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"decart-ai/Lucy-Edit-Dev",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
prompt = ("Change the apron and blouse to a classic clown costume: satin "
|
||||
"polka-dot jumpsuit in bright primary colors, ruffled white collar, "
|
||||
"oversized pom-pom buttons, white gloves, oversized red shoes, red "
|
||||
"foam nose; soft window light from left, eye-level medium shot.")
|
||||
video_path = "https://d2drjpuinn46lb.cloudfront.net/painter_original_edit.mp4"
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
negative_prompt="",
|
||||
video_path=video_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
fps=24,
|
||||
guidance_scale=5.0,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,91 @@
|
||||
"""Run ``judge.third_person_separation`` (needs ``.[eval-judge]`` + a Gemini key)
|
||||
over each baseline and print the candidate's win-rate table — from a ``--manifest``
|
||||
of pairs, or by pairing ``--candidate-dir`` against each ``--reference`` dir by
|
||||
filename stem.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
METRIC = "judge.third_person_separation"
|
||||
VIDEO_EXTS = {".mp4", ".avi", ".mov", ".mkv", ".webm"}
|
||||
IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".webp"}
|
||||
|
||||
|
||||
def _by_stem(directory: Path, exts: set[str]) -> dict[str, Path]:
|
||||
"""Map filename stem -> path for files with the given extensions."""
|
||||
return {p.stem: p for p in sorted(directory.iterdir()) if p.suffix.lower() in exts}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--candidate-dir", type=Path, default=None,
|
||||
help="Directory of candidate clips (directory mode).")
|
||||
p.add_argument("--reference", action="append", default=[], metavar="NAME=DIR",
|
||||
help="Baseline directory, repeatable: 'name=dir' or bare 'dir'.")
|
||||
p.add_argument("--image-dir", type=Path, default=None,
|
||||
help="Optional first-frame images, matched to clips by stem.")
|
||||
p.add_argument("--prompts-json", type=Path, default=None,
|
||||
help="Optional {stem: control-signal text} JSON.")
|
||||
p.add_argument("--actions-json", type=Path, default=None,
|
||||
help="Optional {stem: action-label} JSON for the per-action breakdown.")
|
||||
p.add_argument("--manifest", type=Path, default=None,
|
||||
help="JSON list of {baseline, video_path, reference_path, ...} rows.")
|
||||
p.add_argument("--output", type=Path, default=None)
|
||||
args = p.parse_args()
|
||||
|
||||
# Group path-only samples per baseline: {baseline: [sample dict, ...]}.
|
||||
by_baseline: dict[str, list[dict]] = defaultdict(list)
|
||||
if args.manifest is not None:
|
||||
for row in json.loads(args.manifest.read_text()):
|
||||
by_baseline[row.get("baseline", "baseline")].append(
|
||||
{k: v for k, v in row.items() if k != "baseline"})
|
||||
elif args.candidate_dir is not None and args.reference:
|
||||
cands = _by_stem(args.candidate_dir, VIDEO_EXTS)
|
||||
images = _by_stem(args.image_dir, IMAGE_EXTS) if args.image_dir else {}
|
||||
prompts = json.loads(args.prompts_json.read_text()) if args.prompts_json else {}
|
||||
actions = json.loads(args.actions_json.read_text()) if args.actions_json else {}
|
||||
for spec in args.reference:
|
||||
name, sep, ref_dir = spec.partition("=")
|
||||
if not sep:
|
||||
name, ref_dir = Path(spec).name, spec
|
||||
refs = _by_stem(Path(ref_dir), VIDEO_EXTS)
|
||||
for stem in sorted(cands.keys() & refs.keys()):
|
||||
sample = {"video_path": str(cands[stem]), "reference_path": str(refs[stem])}
|
||||
if stem in images:
|
||||
sample["image_path"] = str(images[stem])
|
||||
if stem in prompts:
|
||||
sample["text_prompt"] = prompts[stem]
|
||||
if stem in actions:
|
||||
sample["action"] = actions[stem]
|
||||
by_baseline[name].append(sample)
|
||||
else:
|
||||
p.error("provide either --manifest, or --candidate-dir with at least one --reference")
|
||||
|
||||
ev = create_evaluator(metrics=[METRIC], device="cpu")
|
||||
print("\n| Baseline | Candidate win-rate (excl. ties) | W / L / T | n |")
|
||||
print("|---|---|---|---|")
|
||||
rows = {}
|
||||
for baseline, samples in by_baseline.items():
|
||||
res = ev.evaluate(samples=samples).corpus[METRIC]
|
||||
rows[baseline] = res
|
||||
d = res.details
|
||||
if res.score is None:
|
||||
print(f"| {baseline} | — | — | 0 |")
|
||||
else:
|
||||
print(f"| {baseline} | {100 * res.score:.1f}% | {d['wins']}/{d['losses']}/{d['ties']} | {d['n']} |")
|
||||
|
||||
if args.output is not None:
|
||||
payload = {b: {"score": r.score, "details": r.details} for b, r in rows.items()}
|
||||
args.output.write_text(json.dumps(payload, indent=2))
|
||||
print(f"\nWrote {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -78,16 +78,26 @@ print(f'{mj}.{mn}')"
|
||||
}
|
||||
|
||||
if [ "${GPU_BACKEND}" = "CUDA" ]; then
|
||||
detected_cc="$(detect_with_torch)" || {
|
||||
echo "ERROR: torch-based CUDA arch detection failed in uv environment." >&2
|
||||
echo " Ensure torch is installed and CUDA is available in the uv-selected Python." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
cc_major="${detected_cc%%.*}"
|
||||
cc_minor="${detected_cc##*.}"
|
||||
# Compute capability drives the arch/TK defaults below. Prefer an explicit
|
||||
# TORCH_CUDA_ARCH_LIST (works on GPU-less build machines such as CI/Docker);
|
||||
# only probe a live GPU via torch when no arch was provided.
|
||||
if [ -n "${TORCH_CUDA_ARCH_LIST:-}" ]; then
|
||||
echo "Using TORCH_CUDA_ARCH_LIST=${TORCH_CUDA_ARCH_LIST} (skipping torch GPU probe)"
|
||||
first_arch="${TORCH_CUDA_ARCH_LIST%%[;, ]*}" # first entry, e.g. 9.0a
|
||||
first_arch="${first_arch%[af]}" # strip trailing a/f suffix
|
||||
cc_major="${first_arch%%.*}"
|
||||
cc_minor="${first_arch##*.}"
|
||||
else
|
||||
detected_cc="$(detect_with_torch)" || {
|
||||
echo "ERROR: torch-based CUDA arch detection failed and TORCH_CUDA_ARCH_LIST is unset." >&2
|
||||
echo " Set TORCH_CUDA_ARCH_LIST (e.g. 9.0a) for GPU-less builds, or build where CUDA is available." >&2
|
||||
exit 1
|
||||
}
|
||||
cc_major="${detected_cc%%.*}"
|
||||
cc_minor="${detected_cc##*.}"
|
||||
echo "Detected compute capability via torch: ${detected_cc}"
|
||||
fi
|
||||
cmake_arch="${cc_major}${cc_minor}"
|
||||
echo "Detected compute capability via torch: ${detected_cc} (sm_${cmake_arch})"
|
||||
|
||||
# Respect explicit overrides.
|
||||
if [ -z "${TORCH_CUDA_ARCH_LIST:-}" ]; then
|
||||
|
||||
@@ -29,6 +29,10 @@ class SamplingParam:
|
||||
# Video inputs
|
||||
video_path: str | None = None
|
||||
|
||||
# Optional pre-generated diffusion latents. Used by parity/debug harnesses
|
||||
# and advanced callers that need deterministic latent reuse.
|
||||
latents: Any | None = None
|
||||
|
||||
# Action control inputs (Matrix-Game)
|
||||
mouse_cond: Any | None = None # Shape: (B, T, 2)
|
||||
keyboard_cond: Any | None = None # Shape: (B, T, K)
|
||||
@@ -64,6 +68,7 @@ class SamplingParam:
|
||||
# Text inputs
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
max_sequence_length: int | None = None
|
||||
prompt_path: str | None = None
|
||||
output_path: str = "outputs/"
|
||||
output_video_name: str | None = None
|
||||
|
||||
@@ -150,14 +150,11 @@ class VideoSparseAttentionMetadata(AttentionMetadata):
|
||||
# in postprocess_output(). Avoids materializing the intermediate
|
||||
# ``[B, len(non_pad_index), H, D]`` tensor on every layer.
|
||||
untile_combined_index: torch.LongTensor
|
||||
# Per-step shared padded buffer used by tile(). Lazily populated on
|
||||
# the first layer's call and reused by every subsequent VSA layer in
|
||||
# the same denoising step. Scoping to metadata (not class/instance)
|
||||
# makes the reuse thread-safe across concurrent requests and keeps
|
||||
# the "pad positions are zero" invariant trivially true (the buffer
|
||||
# is freshly zeroed alongside ``non_pad_index`` so the index set
|
||||
# cannot drift between calls).
|
||||
# Per-step shared padded buffer used by tile(). Inference can reuse this
|
||||
# across VSA layers, but training disables it so activation checkpointing
|
||||
# can release the large tiled QKVG scratch tensor after each attention call.
|
||||
tile_buf: torch.Tensor | None = None
|
||||
cache_tile_buf: bool = True
|
||||
|
||||
|
||||
class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
@@ -175,6 +172,7 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
device: torch.device,
|
||||
cache_tile_buf: bool = True,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> VideoSparseAttentionMetadata:
|
||||
patch_size = patch_size
|
||||
@@ -201,7 +199,8 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
reverse_tile_partition_indices=reverse_tile_partition_indices,
|
||||
variable_block_sizes=variable_block_sizes,
|
||||
non_pad_index=non_pad_index,
|
||||
untile_combined_index=untile_combined_index)
|
||||
untile_combined_index=untile_combined_index,
|
||||
cache_tile_buf=cache_tile_buf)
|
||||
|
||||
|
||||
class VideoSparseAttentionImpl(AttentionImpl):
|
||||
@@ -237,6 +236,11 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
w_padded_size = num_tiles[2] * VSA_TILE_SIZE[2]
|
||||
target_shape = (x.shape[0], t_padded_size * h_padded_size * w_padded_size, x.shape[-2], x.shape[-1])
|
||||
|
||||
if not attn_metadata.cache_tile_buf:
|
||||
buf = torch.zeros(target_shape, device=x.device, dtype=x.dtype)
|
||||
buf[:, attn_metadata.non_pad_index] = x[:, attn_metadata.tile_partition_indices]
|
||||
return buf
|
||||
|
||||
# Reuse the per-step buffer stashed on metadata (lazily allocated
|
||||
# on the first VSA layer's call within a denoising step). Pad
|
||||
# positions are zero from the initial torch.zeros and never
|
||||
|
||||
@@ -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.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -14,5 +15,5 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"MagiHumanVideoConfig", "StableAudioConfig"
|
||||
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
]
|
||||
|
||||
@@ -14,6 +14,11 @@ class DiTArchConfig(ArchConfig):
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
# When True, the denoising stage casts text/prompt embeddings to the DiT's
|
||||
# working dtype before the diffusion loop. Flux2 requires this (BFL casts ctx
|
||||
# to bf16 before denoising); models with fp32 text encoders (Wan, Hunyuan15,
|
||||
# SD3.5) leave it False to preserve full-precision embeddings.
|
||||
cast_prompt_embeds_to_dit_dtype: bool = False
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum,
|
||||
...] = (AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2ArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for Flux2 transformer model."""
|
||||
|
||||
cast_prompt_embeds_to_dit_dtype: bool = True
|
||||
|
||||
# Flux2-specific architecture parameters
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = None
|
||||
num_layers: int = 19 # Number of double-stream transformer blocks
|
||||
num_single_layers: int = 38 # Number of single-stream transformer blocks
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096 # Dimension for text encoder output
|
||||
timestep_guidance_channels: int = 256 # Dimension for timestep embedding
|
||||
mlp_ratio: float = 3.0
|
||||
axes_dims_rope: tuple[int, ...] = (32, 32, 32, 32) # RoPE dimensions per axis (match diffusers Flux2)
|
||||
rope_theta: int = 2000 # Base frequency for RoPE (match diffusers Flux2)
|
||||
eps: float = 1e-6
|
||||
guidance_embeds: bool = True # Whether to use guidance embeddings
|
||||
# When True, compute SwiGLU in fp32 inside ``ff_context`` only (bf16 noise mitigation).
|
||||
ff_context_swiglu_fp32: bool = False
|
||||
|
||||
# Parameter name mapping for loading HuggingFace checkpoints
|
||||
param_names_mapping: dict = field(default_factory=lambda: {
|
||||
r"transformer\.(\w*)\.(.*)$": r"\1.\2",
|
||||
})
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
def update_from_weight_keys(self, all_keys: set[str]) -> None:
|
||||
"""Infer num_layers and num_single_layers from checkpoint weight keys so the model is built with the same number of blocks as the weights."""
|
||||
if not all_keys:
|
||||
return
|
||||
num_layers = 0
|
||||
num_single_layers = 0
|
||||
for k in all_keys:
|
||||
if "single_transformer_blocks." not in k and "transformer_blocks." in k:
|
||||
parts = k.split("transformer_blocks.")[-1].split(".")
|
||||
if parts[0].isdigit():
|
||||
num_layers = max(num_layers, int(parts[0]) + 1)
|
||||
if "single_transformer_blocks." in k:
|
||||
parts = k.split("single_transformer_blocks.")[-1].split(".")
|
||||
if parts[0].isdigit():
|
||||
num_single_layers = max(num_single_layers, int(parts[0]) + 1)
|
||||
if num_layers > 0:
|
||||
self.num_layers = num_layers
|
||||
logger.info("Inferred num_layers=%s from checkpoint keys", num_layers)
|
||||
if num_single_layers > 0:
|
||||
self.num_single_layers = num_single_layers
|
||||
logger.info("Inferred num_single_layers=%s from checkpoint keys", num_single_layers)
|
||||
if num_layers > 0 or num_single_layers > 0:
|
||||
self.__post_init__()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2Config(DiTConfig):
|
||||
"""Configuration for Flux2 transformer model."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=Flux2ArchConfig)
|
||||
|
||||
prefix: str = "Flux"
|
||||
@@ -7,6 +7,8 @@ from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
@@ -15,5 +17,5 @@ __all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig"
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
|
||||
]
|
||||
|
||||
@@ -36,6 +36,11 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
default_factory=list) # mapping from huggingface weight names to custom names
|
||||
tokenizer_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
# When True, the tokenizer loader prefers AutoProcessor over AutoTokenizer
|
||||
# for encoders whose tokenizer dir ships a processor_config.json (e.g. Flux2
|
||||
# full's Mistral3 multimodal processor). Default False keeps every existing
|
||||
# encoder on the historical AutoTokenizer path.
|
||||
require_processor: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Mistral3 text encoder configuration for full Flux2."""
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Mistral3TextArchConfig(TextEncoderArchConfig):
|
||||
"""Architecture config for the Mistral3 text encoder used by full Flux2."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["Mistral3ForConditionalGeneration"])
|
||||
hidden_size: int = 5120
|
||||
num_hidden_layers: int = 40
|
||||
text_len: int = 512
|
||||
output_hidden_states: bool = True
|
||||
# Mistral3 (full Flux2) ships a multimodal processor; load via AutoProcessor.
|
||||
require_processor: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Mistral3TextConfig(TextEncoderConfig):
|
||||
"""Top-level config for the Mistral3 full Flux2 text encoder."""
|
||||
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=Mistral3TextArchConfig)
|
||||
prefix: str = "mistral3"
|
||||
is_chat_model: bool = True
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Ported from SGLang: python/sglang/multimodal_gen/configs/models/encoders/qwen3.py
|
||||
"""Qwen3 text encoder configuration for FastVideo diffusion models (e.g. Flux2 Klein)."""
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m: Any) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m: Any) -> bool:
|
||||
return n.endswith("embed_tokens")
|
||||
|
||||
|
||||
def _is_final_norm(n: str, m: Any) -> bool:
|
||||
return n.endswith("norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen3TextArchConfig(TextEncoderArchConfig):
|
||||
"""Architecture config for Qwen3 text encoder.
|
||||
|
||||
Qwen3 is similar to LLaMA but with QK-Norm (RMSNorm on Q and K before attention).
|
||||
Used by Flux2 Klein.
|
||||
"""
|
||||
|
||||
vocab_size: int = 151936
|
||||
hidden_size: int = 2560
|
||||
intermediate_size: int = 9728
|
||||
num_hidden_layers: int = 36
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 8
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 40960
|
||||
initializer_range: float = 0.02
|
||||
rms_norm_eps: float = 1e-6
|
||||
use_cache: bool = True
|
||||
pad_token_id: int = 151643
|
||||
bos_token_id: int = 151643
|
||||
eos_token_id: int = 151645
|
||||
tie_word_embeddings: bool = True
|
||||
rope_theta: float = 1000000.0
|
||||
rope_scaling: dict | None = None
|
||||
attention_bias: bool = False
|
||||
attention_dropout: float = 0.0
|
||||
mlp_bias: bool = False
|
||||
head_dim: int = 128
|
||||
text_len: int = 512
|
||||
output_hidden_states: bool = True # Klein needs hidden states from layers 9, 18, 27
|
||||
|
||||
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=lambda: [
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen3TextConfig(TextEncoderConfig):
|
||||
"""Top-level config for Qwen3 text encoder."""
|
||||
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=Qwen3TextArchConfig)
|
||||
prefix: str = "qwen3"
|
||||
is_chat_model: bool = True
|
||||
@@ -6,6 +6,7 @@ from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
@@ -19,4 +20,5 @@ __all__ = [
|
||||
"LTX2VAEConfig",
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2VAEArchConfig(VAEArchConfig):
|
||||
"""Architecture configuration for Flux2 VAE model."""
|
||||
|
||||
# Flux2 VAE-specific architecture parameters
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
down_block_types: tuple[str, ...] = (
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"AttnDownEncoderBlock2D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
"AttnUpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
)
|
||||
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
|
||||
layers_per_block: int = 2
|
||||
act_fn: str = "silu"
|
||||
latent_channels: int = 16
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 512
|
||||
force_upcast: bool = False
|
||||
use_quant_conv: bool = True
|
||||
use_post_quant_conv: bool = True
|
||||
mid_block_add_attention: bool = True
|
||||
batch_norm_eps: float = 1e-5
|
||||
batch_norm_momentum: float = 0.1
|
||||
patch_size: tuple[int, int] = (1, 1)
|
||||
|
||||
# Latent scaling for decode: avoid division-by-zero; match Flux/Flux2 convention (e.g. 0.13025)
|
||||
scaling_factor: float = 0.13025
|
||||
|
||||
# Spatial compression (for images, this is typically 8)
|
||||
spatial_compression_ratio: int = 8
|
||||
temporal_compression_ratio: int = 1 # Images don't have temporal dimension
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2VAEConfig(VAEConfig):
|
||||
"""Configuration for Flux2 VAE model."""
|
||||
|
||||
arch_config: Flux2VAEArchConfig = field(default_factory=Flux2VAEArchConfig)
|
||||
|
||||
# Flux2 is an image model, so disable temporal tiling
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
@@ -9,12 +9,12 @@ from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
from fastvideo.configs.pipelines.wan import (LucyEditDevConfig, SelfForcingWanT2V480PConfig, WanI2V480PConfig,
|
||||
WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
|
||||
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -35,6 +35,11 @@ class PipelineConfig:
|
||||
flow_shift: float | None = None
|
||||
flow_shift_sr: float | None = None
|
||||
disable_autocast: bool = False
|
||||
# When True, the scheduler's Euler update runs in fp32 outside the autocast
|
||||
# block (Diffusers-style; avoids BF16 drift over multiple steps). Flux2 sets
|
||||
# this True for reference parity; other models keep the legacy in-autocast
|
||||
# behavior to preserve existing SSIM references.
|
||||
scheduler_step_in_fp32: bool = False
|
||||
is_causal: bool = False
|
||||
|
||||
# Model configuration
|
||||
@@ -64,8 +69,9 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
# Wan2.2 task modifiers
|
||||
ti2v_task: bool = False
|
||||
lucy_edit_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
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.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.base import EncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2PipelineConfig(PipelineConfig):
|
||||
"""Configuration for Flux2 image generation pipeline."""
|
||||
|
||||
# Flux2-specific parameters
|
||||
embedded_cfg_scale: float | None = 4.0
|
||||
scheduler_step_in_fp32: bool = True
|
||||
flux2_text_encoder_type: str = "mistral3"
|
||||
text_encoder_out_layers: tuple[int, ...] = (10, 20, 30)
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=Flux2Config)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=Flux2VAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = False # Flux2 is image model, disable tiling by default
|
||||
vae_sp: bool = False
|
||||
|
||||
# Text encoder configuration (full Flux2 uses Mistral3)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Mistral3TextConfig(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Default postprocess function (can be overridden)
|
||||
@staticmethod
|
||||
def default_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Default text postprocessing for Flux2."""
|
||||
return outputs.last_hidden_state
|
||||
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (Flux2PipelineConfig.default_postprocess_text, ))
|
||||
|
||||
|
||||
def flux2_klein_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Klein postprocess: hidden states from layers 9, 18, 27 (Qwen3)."""
|
||||
hidden_states_layers: list[int] = [9, 18, 27]
|
||||
if outputs.hidden_states is None:
|
||||
raise ValueError("Flux2 Klein requires output_hidden_states=True from text encoder")
|
||||
out = torch.stack([outputs.hidden_states[k] for k in hidden_states_layers], dim=1)
|
||||
batch_size, num_channels, seq_len, hidden_dim = out.shape
|
||||
prompt_embeds = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, num_channels * hidden_dim)
|
||||
return prompt_embeds
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2KleinEncoderArchConfig(EncoderArchConfig):
|
||||
"""Encoder arch config for Flux2 Klein (Qwen3); needs hidden states for layers 9, 18, 27."""
|
||||
output_hidden_states: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2KleinTextEncoderConfig(EncoderConfig):
|
||||
"""Text encoder config for Flux2 Klein (Qwen3)."""
|
||||
arch_config: EncoderArchConfig = field(default_factory=Flux2KleinEncoderArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Flux2KleinPipelineConfig(Flux2PipelineConfig):
|
||||
"""Configuration for Flux2 Klein (distilled, 4-step, no guidance)."""
|
||||
embedded_cfg_scale: float | None = None # Klein distilled: no guidance embedding (matches Diffusers)
|
||||
scheduler_step_in_fp32: bool = True
|
||||
flux2_text_encoder_type: str = "qwen3"
|
||||
text_encoder_out_layers: tuple[int, ...] = (9, 18, 27)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Qwen3TextConfig(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (flux2_klein_postprocess_text, ))
|
||||
@@ -6,9 +6,11 @@ import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput, CLIPVisionConfig, T5Config,
|
||||
WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@@ -120,6 +122,142 @@ class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
expand_timesteps: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
assert not (self.ti2v_task and self.lucy_edit_task)
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
class LucyEditDevConfig(Wan2_2_TI2V_5B_Config):
|
||||
"""Configuration for Decart Lucy Edit Dev video editing."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=lambda: WanVideoConfig(arch_config=WanVideoArchConfig(
|
||||
num_attention_heads=24,
|
||||
in_channels=96,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
)))
|
||||
vae_config: VAEConfig = field(default_factory=lambda: WanVAEConfig(arch_config=WanVAEArchConfig(
|
||||
base_dim=160,
|
||||
decoder_base_dim=256,
|
||||
z_dim=48,
|
||||
in_channels=12,
|
||||
out_channels=12,
|
||||
scale_factor_spatial=16,
|
||||
patch_size=2,
|
||||
is_residual=True,
|
||||
clip_output=False,
|
||||
latents_mean=(
|
||||
-0.2289,
|
||||
-0.0052,
|
||||
-0.1323,
|
||||
-0.2339,
|
||||
-0.2799,
|
||||
0.0174,
|
||||
0.1838,
|
||||
0.1557,
|
||||
-0.1382,
|
||||
0.0542,
|
||||
0.2813,
|
||||
0.0891,
|
||||
0.1570,
|
||||
-0.0098,
|
||||
0.0375,
|
||||
-0.1825,
|
||||
-0.2246,
|
||||
-0.1207,
|
||||
-0.0698,
|
||||
0.5109,
|
||||
0.2665,
|
||||
-0.2108,
|
||||
-0.2158,
|
||||
0.2502,
|
||||
-0.2055,
|
||||
-0.0322,
|
||||
0.1109,
|
||||
0.1567,
|
||||
-0.0729,
|
||||
0.0899,
|
||||
-0.2799,
|
||||
-0.1230,
|
||||
-0.0313,
|
||||
-0.1649,
|
||||
0.0117,
|
||||
0.0723,
|
||||
-0.2839,
|
||||
-0.2083,
|
||||
-0.0520,
|
||||
0.3748,
|
||||
0.0152,
|
||||
0.1957,
|
||||
0.1433,
|
||||
-0.2944,
|
||||
0.3573,
|
||||
-0.0548,
|
||||
-0.1681,
|
||||
-0.0667,
|
||||
),
|
||||
latents_std=(
|
||||
0.4765,
|
||||
1.0364,
|
||||
0.4514,
|
||||
1.1677,
|
||||
0.5313,
|
||||
0.4990,
|
||||
0.4818,
|
||||
0.5013,
|
||||
0.8158,
|
||||
1.0344,
|
||||
0.5894,
|
||||
1.0901,
|
||||
0.6885,
|
||||
0.6165,
|
||||
0.8454,
|
||||
0.4978,
|
||||
0.5759,
|
||||
0.3523,
|
||||
0.7135,
|
||||
0.6804,
|
||||
0.5833,
|
||||
1.4146,
|
||||
0.8986,
|
||||
0.5659,
|
||||
0.7069,
|
||||
0.5338,
|
||||
0.4889,
|
||||
0.4917,
|
||||
0.4069,
|
||||
0.4999,
|
||||
0.6866,
|
||||
0.4093,
|
||||
0.5709,
|
||||
0.6065,
|
||||
0.6415,
|
||||
0.4944,
|
||||
0.5726,
|
||||
1.2042,
|
||||
0.5458,
|
||||
1.6887,
|
||||
0.3971,
|
||||
1.0600,
|
||||
0.3943,
|
||||
0.5537,
|
||||
0.5444,
|
||||
0.4089,
|
||||
0.7468,
|
||||
0.7744,
|
||||
),
|
||||
)))
|
||||
ti2v_task: bool = False
|
||||
lucy_edit_task: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
assert not (self.ti2v_task and self.lucy_edit_task)
|
||||
# Lucy uses Wan2.2's enhanced 48-channel VAE latents. Denoising
|
||||
# concatenates noise + video latents, matching the 96-channel
|
||||
# transformer input declared above.
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
@@ -517,6 +517,7 @@ class VideoGenerator:
|
||||
if _ek in kwargs:
|
||||
extra_overrides[_ek] = kwargs.pop(_ek)
|
||||
|
||||
prompt_embeds = kwargs.pop("prompt_embeds", None)
|
||||
sampling_param.update(kwargs)
|
||||
kwargs["_extra_overrides"] = extra_overrides
|
||||
|
||||
@@ -567,6 +568,8 @@ class VideoGenerator:
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
output_path = self._prepare_output_path(sampling_param.output_path, prompt)
|
||||
kwargs["output_path"] = output_path
|
||||
if prompt_embeds is not None:
|
||||
kwargs["prompt_embeds"] = prompt_embeds
|
||||
return self._generate_single_video(
|
||||
prompt=prompt,
|
||||
sampling_param=sampling_param,
|
||||
@@ -669,6 +672,7 @@ class VideoGenerator:
|
||||
prompt = prompt.strip()
|
||||
sampling_param = deepcopy(sampling_param)
|
||||
output_path = kwargs["output_path"]
|
||||
prompt_embeds = kwargs.get("prompt_embeds")
|
||||
sampling_param.prompt = prompt
|
||||
# Process negative prompt
|
||||
if sampling_param.negative_prompt is not None:
|
||||
@@ -715,6 +719,9 @@ class VideoGenerator:
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
)
|
||||
# Allow precomputed prompt_embeds (e.g. from diffusers) to skip text encoding
|
||||
if prompt_embeds is not None:
|
||||
batch.prompt_embeds = (list(prompt_embeds) if isinstance(prompt_embeds, list | tuple) else [prompt_embeds])
|
||||
|
||||
extra_overrides = kwargs.pop("_extra_overrides", {})
|
||||
for _ek, _ev in extra_overrides.items():
|
||||
|
||||
@@ -2,8 +2,9 @@
|
||||
|
||||
In-process evaluation suite for video generations. Includes pixel
|
||||
metrics (SSIM, PSNR, LPIPS), Fréchet Video Distance (FVD), optical-flow
|
||||
comparisons, the full VBench suite, Physics-IQ, audio metrics, and a
|
||||
VLM scorer behind a single registry-driven API.
|
||||
comparisons, the full VBench suite, Physics-IQ, audio metrics, an
|
||||
absolute VLM scorer (`videoscore2`), and a pairwise VLM judge
|
||||
(`judge.third_person_separation`) — all behind a single registry-driven API.
|
||||
|
||||
## Install
|
||||
|
||||
@@ -159,6 +160,7 @@ fastvideo/
|
||||
│ ├── audio/ # clap_score, audiobox_aesthetics, kl_divergence,
|
||||
│ │ # frechet_distance, wer, desync, imagebind_score
|
||||
│ ├── videoscore2/ # VideoScore-2 (Qwen2.5-VL)
|
||||
│ ├── judge/ # pairwise VLM judges (third_person_separation)
|
||||
│ ├── physics_iq/ # PhysicsIQ + sub-metrics
|
||||
│ └── vbench/ # adapter: sys.path bootstrap + shims
|
||||
│ ├── __init__.py
|
||||
@@ -313,6 +315,40 @@ to control read/write behavior. The example script
|
||||
`examples/inference/eval/eval_fvd.py` demonstrates the full
|
||||
two-directory workflow.
|
||||
|
||||
## `judge.third_person_separation` — pairwise VLM judge
|
||||
|
||||
A **preference** metric (a judge, not an absolute score), and the suite's first
|
||||
remote-API one. For each pair the judge (Gemini) sees the shared first frame and
|
||||
two rollouts — a candidate and a reference model under the same control signal —
|
||||
and picks the one that better separates the third-person CHARACTER (foreground)
|
||||
from the BACKGROUND. The corpus score is the candidate's win-rate, excluding
|
||||
ties. Set-vs-set, motion-first; it reads native mp4s, so samples carry path
|
||||
strings, not decoded tensors.
|
||||
|
||||
```bash
|
||||
uv pip install -e .[eval-judge] # opt-in: needs network + an API key
|
||||
export GEMINI_API_KEY=... # or GOOGLE_API_KEY, or ~/.gemini_token
|
||||
```
|
||||
|
||||
```python
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
ev = create_evaluator(metrics=["judge.third_person_separation"], device="cpu")
|
||||
result = ev.evaluate(samples=[
|
||||
{"video_path": "cand/000.mp4", "reference_path": "base/000.mp4",
|
||||
"image_path": "frames/000.png", "text_prompt": "W: moves forward", "action": "W"},
|
||||
# ... more pairs ...
|
||||
]).corpus["judge.third_person_separation"]
|
||||
result.score # candidate win-rate excl. ties; result.details has the breakdown
|
||||
```
|
||||
|
||||
Only `video_path`/`reference_path` are required; `image_path`/`text_prompt`/
|
||||
`action` are optional. Verdicts are cached under `${FASTVIDEO_EVAL_CACHE}/eval/judge/`.
|
||||
The judge separates best when the control yields genuine parallax (e.g.
|
||||
translation); rigid whole-frame motion (e.g. pure camera rotation) is harder. To
|
||||
sweep several baselines into a table, see
|
||||
`examples/inference/eval/eval_third_person_separation.py`.
|
||||
|
||||
## Out of scope (follow-up PRs)
|
||||
|
||||
- **MIND** metrics. Depend on a separate `vipe` upstream submodule.
|
||||
|
||||
@@ -0,0 +1,306 @@
|
||||
"""Pairwise VLM judge (Gemini): of two rollouts under the same control, which
|
||||
better separates the third-person character (foreground) from the background;
|
||||
the score is the candidate's win-rate over a reference.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.models import get_cache_dir
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult, Video
|
||||
|
||||
# Part of the on-disk cache key; bump to invalidate cached verdicts.
|
||||
RUBRIC_ID = "v1"
|
||||
DEFAULT_MODEL = "gemini-2.5-pro"
|
||||
DEFAULT_K = 3
|
||||
|
||||
SYSTEM_PROMPT = ("You are a strict comparative evaluator of third-person video-game "
|
||||
"rollouts. Two videos were generated by two different models from the "
|
||||
"SAME first frame and the SAME control signal. You judge which video "
|
||||
"better demonstrates that the model SEPARATES the third-person CHARACTER "
|
||||
"(foreground) from the BACKGROUND SCENE — i.e. the control animates the "
|
||||
"character as an INDEPENDENT AGENT with its own trajectory while the "
|
||||
"background moves with the camera. You do NOT reward whichever video "
|
||||
"merely looks cleaner, sharper, or higher-res.")
|
||||
|
||||
RUBRIC = """\
|
||||
You are watching TWO generated video rollouts (Video 1, Video 2) played in full,
|
||||
from the same first frame under the same control signal. Pick the better
|
||||
third-person world-model rollout. Judge THREE things together, in this order:
|
||||
|
||||
(C) MOTION / ACTION EXECUTION FIRST. The rollout must actually CARRY OUT the
|
||||
control signal with substantial motion (the scene/character clearly moves as
|
||||
commanded). A clip that is near-static, barely drifts, or only twitches has
|
||||
NOT demonstrated controllable separation — it FAILS, no matter how clean it
|
||||
looks. CRITICAL: do NOT reward a clip for looking "smoother" or "more
|
||||
stable" when that smoothness is really just the ABSENCE OF MOTION. Less
|
||||
motion is NOT better. If one clip executes the action with clear motion and
|
||||
the other is comparatively static, the MOVING one wins (unless it fails B).
|
||||
|
||||
(A) TEMPORAL COHERENCE — among clips that actually move, penalize GENUINE
|
||||
corruption: flicker/strobing, texture boiling/crawling, geometry swimming,
|
||||
the character or scene morphing/warping into mush, colors pulsing, or
|
||||
progressive degradation into noise. Do NOT confuse LEGITIMATE large motion
|
||||
(camera sweeping, character running, scene flowing past) with instability —
|
||||
fast correct motion is GOOD, not a defect. Only true frame-to-frame
|
||||
INCOHERENCE counts against a clip.
|
||||
|
||||
(B) FOREGROUND/BACKGROUND SEPARATION — the character stays a distinct, coherent
|
||||
entity with its own trajectory while the background responds to the control;
|
||||
it does not dissolve/smear into the bg, and the whole frame does not slide
|
||||
as one rigid sheet.
|
||||
|
||||
Decision: among clips that genuinely execute the motion (C), pick the one that
|
||||
is both temporally coherent (A) and shows cleaner separation (B). A static or
|
||||
barely-moving clip loses to a moving one. Real motion is not instability. Do not
|
||||
reward resolution or placidity. "tie" only if truly equivalent on all three."""
|
||||
|
||||
USER_TASK = ("Output a JSON object with four fields:\n"
|
||||
" - video_1_analysis: FIRST, how much does Video 1 actually move — does it "
|
||||
"execute the control with clear motion, or is it near-static / barely "
|
||||
"drifting? THEN: among its motion, is there GENUINE corruption (flicker, "
|
||||
"boiling, morphing into mush) as opposed to legitimate fast motion? THEN: "
|
||||
"is the character a distinct coherent entity vs rigid-slide / dissolve?\n"
|
||||
" - video_2_analysis: the same three checks for Video 2.\n"
|
||||
" - comparison: apply (C) motion-first, then (A) genuine-coherence, then "
|
||||
"(B) separation. A near-static clip loses to a moving one; legitimate large "
|
||||
"motion is NOT a defect.\n"
|
||||
" - winner: \"video_1\", \"video_2\", or \"tie\".")
|
||||
|
||||
RESPONSE_SCHEMA = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"video_1_analysis": {
|
||||
"type": "string"
|
||||
},
|
||||
"video_2_analysis": {
|
||||
"type": "string"
|
||||
},
|
||||
"comparison": {
|
||||
"type": "string"
|
||||
},
|
||||
"winner": {
|
||||
"type": "string",
|
||||
"enum": ["video_1", "video_2", "tie"]
|
||||
},
|
||||
},
|
||||
"required": ["video_1_analysis", "video_2_analysis", "comparison", "winner"],
|
||||
"propertyOrdering": ["video_1_analysis", "video_2_analysis", "comparison", "winner"],
|
||||
}
|
||||
|
||||
|
||||
def _path_of(sample: dict, key: str) -> str | None:
|
||||
"""Resolve a native-file path from a string key or a Video wrapper."""
|
||||
p = sample.get(f"{key}_path")
|
||||
if isinstance(p, str):
|
||||
return p
|
||||
v = sample.get(key)
|
||||
if isinstance(v, Video) and isinstance(v.source, str):
|
||||
return v.source
|
||||
if isinstance(v, str):
|
||||
return v
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_api_key() -> str:
|
||||
for env in ("GEMINI_API_KEY", "GOOGLE_API_KEY"):
|
||||
key = os.environ.get(env)
|
||||
if key:
|
||||
return key.strip()
|
||||
token = Path("~/.gemini_token").expanduser()
|
||||
if token.is_file():
|
||||
return token.read_text().strip()
|
||||
raise ValueError("judge.third_person_separation needs a Gemini API key. Set "
|
||||
"GEMINI_API_KEY (or GOOGLE_API_KEY), or write it to ~/.gemini_token.")
|
||||
|
||||
|
||||
@register("judge.third_person_separation")
|
||||
class ThirdPersonSeparationMetric(BaseMetric):
|
||||
"""Pairwise VLM judge of third-person fg/bg separation; corpus win-rate."""
|
||||
|
||||
name = "judge.third_person_separation"
|
||||
requires_reference = True
|
||||
higher_is_better = True
|
||||
needs_gpu = False
|
||||
is_set_metric = True
|
||||
dependencies = ["google.genai"]
|
||||
|
||||
def __init__(self, model: str = DEFAULT_MODEL, k: int = DEFAULT_K) -> None:
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.k = k
|
||||
self._client: Any = None
|
||||
self._files: dict[str, Any] = {} # path -> uploaded Gemini file handle
|
||||
self._records: list[dict] = [] # one per accumulated pair
|
||||
|
||||
# --- model / client -----------------------------------------------------
|
||||
def setup(self) -> None:
|
||||
if self._client is not None:
|
||||
return
|
||||
from google import genai
|
||||
self._client = genai.Client(api_key=_resolve_api_key())
|
||||
|
||||
# --- set-vs-set protocol ------------------------------------------------
|
||||
def reset(self) -> None:
|
||||
self._records = []
|
||||
self._files = {}
|
||||
|
||||
def accumulate(self, sample: dict) -> None:
|
||||
cand = _path_of(sample, "video")
|
||||
base = _path_of(sample, "reference")
|
||||
if cand is None or base is None:
|
||||
return # nothing to compare
|
||||
|
||||
image = sample.get("image_path")
|
||||
action_text = sample.get("text_prompt") or ""
|
||||
action = sample.get("action")
|
||||
|
||||
rec = self._cached(cand, base, action_text)
|
||||
if rec is None:
|
||||
if self._client is None:
|
||||
self.setup()
|
||||
rec = self._judge_pair(cand, base, image, action_text)
|
||||
self._write_cache(cand, base, action_text, rec)
|
||||
rec = {**rec, "action": action}
|
||||
self._records.append(rec)
|
||||
|
||||
def finalize(self) -> MetricResult:
|
||||
recs = [r for r in self._records if r.get("verdict")]
|
||||
if not recs:
|
||||
return MetricResult(name=self.name, score=None, details={"skipped": "no pairs judged"})
|
||||
wins = sum(r["verdict"] == "candidate" for r in recs)
|
||||
losses = sum(r["verdict"] == "baseline" for r in recs)
|
||||
ties = sum(r["verdict"] == "tie" for r in recs)
|
||||
decided = wins + losses
|
||||
score = wins / decided if decided else None
|
||||
|
||||
# Per-action win-rate, grouped by the raw label (no assumed control scheme).
|
||||
per_action: dict[str, dict] = {}
|
||||
labels: set[str] = {str(r["action"]) for r in recs if r.get("action")}
|
||||
for action in sorted(labels):
|
||||
gr = [r for r in recs if r.get("action") == action]
|
||||
gw = sum(r["verdict"] == "candidate" for r in gr)
|
||||
gl = sum(r["verdict"] == "baseline" for r in gr)
|
||||
per_action[action] = {
|
||||
"n": len(gr),
|
||||
"wins": gw,
|
||||
"losses": gl,
|
||||
"ties": len(gr) - gw - gl,
|
||||
"win_rate": (gw / (gw + gl)) if (gw + gl) else None,
|
||||
}
|
||||
return MetricResult(name=self.name,
|
||||
score=score,
|
||||
details={
|
||||
"wins": wins,
|
||||
"losses": losses,
|
||||
"ties": ties,
|
||||
"n": len(recs),
|
||||
"win_rate_excl_ties": score,
|
||||
"per_action": per_action,
|
||||
})
|
||||
|
||||
def merge_from(self, other: BaseMetric) -> None:
|
||||
assert isinstance(other, ThirdPersonSeparationMetric)
|
||||
self._records.extend(other._records)
|
||||
|
||||
# --- judging ------------------------------------------------------------
|
||||
def _judge_pair(self, cand: str, base: str, image: str | None, action_text: str) -> dict:
|
||||
"""k counterbalanced comparisons → aggregated per-pair verdict."""
|
||||
# Seed the A/B alternation from the pair itself so it is reproducible and
|
||||
# independent of evaluation order or which subset is being run.
|
||||
seed = int(hashlib.sha1(f"{cand}|{base}".encode()).hexdigest(), 16)
|
||||
mapped: list[str] = []
|
||||
for i in range(self.k):
|
||||
cand_first = (seed + i) % 2 == 0
|
||||
v1, v2 = (cand, base) if cand_first else (base, cand)
|
||||
winner = self._one_call(image, v1, v2, action_text)
|
||||
if winner == "tie":
|
||||
mapped.append("tie")
|
||||
elif (winner == "video_1") == cand_first:
|
||||
mapped.append("candidate")
|
||||
else:
|
||||
mapped.append("baseline")
|
||||
cand_w = mapped.count("candidate")
|
||||
base_w = mapped.count("baseline")
|
||||
verdict = ("candidate" if cand_w > base_w else "baseline" if base_w > cand_w else "tie")
|
||||
return {
|
||||
"verdict": verdict,
|
||||
"candidate_wins": cand_w,
|
||||
"baseline_wins": base_w,
|
||||
"ties": mapped.count("tie"),
|
||||
"k": self.k,
|
||||
"rubric_id": RUBRIC_ID
|
||||
}
|
||||
|
||||
def _one_call(self, image: str | None, vid1: str, vid2: str, action_text: str) -> str:
|
||||
from google.genai import types
|
||||
contents: list[Any] = [action_text or "Compare these two rollouts."]
|
||||
if image is not None:
|
||||
contents += ["\nFirst frame (input condition, shared by BOTH "
|
||||
"videos):", self._upload(image)]
|
||||
contents += [
|
||||
"\nVideo 1 (model A's full rollout — watch it in motion):",
|
||||
self._upload(vid1),
|
||||
"\nVideo 2 (model B's full rollout — watch it in motion):",
|
||||
self._upload(vid2),
|
||||
"\n" + RUBRIC + "\n\n" + USER_TASK,
|
||||
]
|
||||
for attempt in range(6):
|
||||
try:
|
||||
resp = self._client.models.generate_content(model=self.model,
|
||||
contents=contents,
|
||||
config=types.GenerateContentConfig(
|
||||
system_instruction=SYSTEM_PROMPT,
|
||||
response_mime_type="application/json",
|
||||
response_schema=RESPONSE_SCHEMA,
|
||||
temperature=0.4))
|
||||
return json.loads(resp.text).get("winner", "tie")
|
||||
except Exception as exc: # noqa: BLE001 - transient API errors
|
||||
if attempt == 5:
|
||||
print(f"[judge] giving up after 6 attempts ({exc}); scoring this call a tie")
|
||||
break
|
||||
is_429 = "429" in str(exc) or "RESOURCE_EXHAUSTED" in str(exc)
|
||||
time.sleep(40 if is_429 else 2**attempt)
|
||||
return "tie"
|
||||
|
||||
def _upload(self, path: str) -> Any:
|
||||
f = self._files.get(path)
|
||||
if f is not None:
|
||||
return f
|
||||
f = self._client.files.upload(file=path)
|
||||
while f.state.name != "ACTIVE":
|
||||
time.sleep(1)
|
||||
f = self._client.files.get(name=f.name)
|
||||
if f.state.name == "FAILED":
|
||||
raise RuntimeError(f"Gemini upload failed for {path}")
|
||||
self._files[path] = f
|
||||
return f
|
||||
|
||||
# --- per-pair cache -----------------------------------------------------
|
||||
def _cache_path(self, cand: str, base: str, action_text: str) -> Path:
|
||||
# k is intentionally NOT in the key so a larger-k run can reuse an
|
||||
# existing verdict with enough samples (see ``_cached``).
|
||||
h = hashlib.sha1(f"{RUBRIC_ID}|{self.model}|{cand}|{base}|{action_text}".encode()).hexdigest()[:16]
|
||||
return get_cache_dir() / "judge" / "third_person_separation" / f"{h}.json"
|
||||
|
||||
def _cached(self, cand: str, base: str, action_text: str) -> dict | None:
|
||||
cp = self._cache_path(cand, base, action_text)
|
||||
if not cp.is_file():
|
||||
return None
|
||||
try:
|
||||
rec = json.loads(cp.read_text())
|
||||
except Exception:
|
||||
return None
|
||||
return rec if rec.get("k", 0) >= self.k and "verdict" in rec else None
|
||||
|
||||
def _write_cache(self, cand: str, base: str, action_text: str, rec: dict) -> None:
|
||||
cp = self._cache_path(cand, base, action_text)
|
||||
cp.parent.mkdir(parents=True, exist_ok=True)
|
||||
cp.write_text(json.dumps(rec, indent=2))
|
||||
@@ -85,6 +85,8 @@ def _extra_for(metric_name: str) -> str:
|
||||
return "eval-physics-iq"
|
||||
if metric_name.startswith("audio."):
|
||||
return "eval-audio"
|
||||
if metric_name.startswith("judge."):
|
||||
return "eval-judge"
|
||||
return "eval"
|
||||
|
||||
|
||||
|
||||
@@ -219,6 +219,15 @@ class ReplicatedLinear(LinearBase):
|
||||
(e.g. model.layers.0.qkv_proj)
|
||||
"""
|
||||
|
||||
# Opt-in instrumentation: when ``enable_shape_tracking`` is set to True,
|
||||
# ``forward`` records every unique ``(input_shape, output_shape)`` pair
|
||||
# observed across all ``ReplicatedLinear`` instances, along with the
|
||||
# subclass name that produced it. Used by upcoming QAT-aware backends
|
||||
# to discover which GEMM shapes need quantized kernels. Defaults to
|
||||
# False; default forward path is bit-identical to pre-slice behavior.
|
||||
enable_shape_tracking = False
|
||||
_shape_to_layer_types: dict[tuple[torch.Size, torch.Size], set[str]] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
@@ -285,6 +294,8 @@ class ReplicatedLinear(LinearBase):
|
||||
bias = self.bias if not self.skip_bias_add else None
|
||||
assert self.quant_method is not None
|
||||
output = self.quant_method.apply(self, x, bias)
|
||||
if self.enable_shape_tracking:
|
||||
self._track_shape(x.shape, output.shape)
|
||||
output_bias = self.bias if self.skip_bias_add else None
|
||||
return output, output_bias
|
||||
|
||||
@@ -294,6 +305,41 @@ class ReplicatedLinear(LinearBase):
|
||||
s += f", bias={self.bias is not None}"
|
||||
return s
|
||||
|
||||
@classmethod
|
||||
def get_shape_mapping(cls) -> dict:
|
||||
"""Get the mapping from (input_shape, output_shape) to layer types."""
|
||||
return cls._shape_to_layer_types.copy()
|
||||
|
||||
@classmethod
|
||||
def reset_shape_tracking(cls) -> None:
|
||||
"""Clear tracked shapes and layer type mappings."""
|
||||
cls._shape_to_layer_types.clear()
|
||||
|
||||
def _track_shape(self, input_shape: torch.Size, output_shape: torch.Size) -> None:
|
||||
shape_key = (input_shape, output_shape)
|
||||
if shape_key not in self._shape_to_layer_types:
|
||||
self._shape_to_layer_types[shape_key] = set()
|
||||
logger.debug("Layer: %s | input shape: %s --> output shape: %s, Quant Method: %s", self.prefix, input_shape,
|
||||
output_shape, self.quant_method.__class__.__name__)
|
||||
self._shape_to_layer_types[shape_key].add(self.__class__.__name__)
|
||||
|
||||
@classmethod
|
||||
def print_shape_summary(cls) -> None:
|
||||
"""Log a summary of all unique shapes and their layer types."""
|
||||
if not cls._shape_to_layer_types:
|
||||
logger.info("No shapes have been processed yet.")
|
||||
return
|
||||
|
||||
lines = [
|
||||
"=== Matrix Multiplication Shape Summary ===",
|
||||
f"Total unique shapes: {len(cls._shape_to_layer_types)}",
|
||||
]
|
||||
for i, (shape_key, layer_types) in enumerate(cls._shape_to_layer_types.items(), 1):
|
||||
input_shape, output_shape = shape_key
|
||||
lines.append(f"{i}. Input: {input_shape} → Output: {output_shape}")
|
||||
lines.append(f" Layer types: {', '.join(sorted(layer_types))}")
|
||||
logger.info("\n".join(lines))
|
||||
|
||||
|
||||
class ColumnParallelLinear(LinearBase):
|
||||
"""Linear layer with column parallelism.
|
||||
|
||||
+12
-2
@@ -5,6 +5,7 @@ import torch.nn as nn
|
||||
|
||||
from fastvideo.layers.activation import get_act_fn
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
@@ -21,18 +22,27 @@ class MLP(nn.Module):
|
||||
act_type: str = "gelu_pytorch_tanh",
|
||||
dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
input_dim,
|
||||
mlp_hidden_dim, # For activation func like SiLU that need 2x width
|
||||
bias=bias,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc_in",
|
||||
)
|
||||
|
||||
self.act = get_act_fn(act_type)
|
||||
if output_dim is None:
|
||||
output_dim = input_dim
|
||||
self.fc_out = ReplicatedLinear(mlp_hidden_dim, output_dim, bias=bias, params_dtype=dtype)
|
||||
self.fc_out = ReplicatedLinear(mlp_hidden_dim,
|
||||
output_dim,
|
||||
bias=bias,
|
||||
params_dtype=dtype,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc_out")
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x, _ = self.fc_in(x)
|
||||
|
||||
@@ -52,6 +52,7 @@ def apply_rotary_emb(
|
||||
freqs_cis: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
|
||||
use_real: bool = True,
|
||||
use_real_unbind_dim: int = -1,
|
||||
sequence_dim: int = 2,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
|
||||
@@ -60,17 +61,23 @@ def apply_rotary_emb(
|
||||
tensors contain rotary embeddings and are returned as real tensors.
|
||||
Args:
|
||||
x (`torch.Tensor`):
|
||||
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
|
||||
Query or key tensor to apply rotary embeddings. [B, H, S, D] if sequence_dim=2 else [B, S, H, D].
|
||||
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
|
||||
sequence_dim: 1 = sequence at dim 1 (cos [1,S,1,D], x [B,S,H,D]); 2 = sequence at dim 2 (cos [1,1,S,D], x [B,H,S,D]).
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
||||
"""
|
||||
if use_real:
|
||||
cos, sin = freqs_cis # [S, D]
|
||||
# Match Diffusers broadcasting (sequence_dim=2 case)
|
||||
cos = cos[None, None, :, :]
|
||||
sin = sin[None, None, :, :]
|
||||
cos, sin = cos.to(x.device), sin.to(x.device)
|
||||
if sequence_dim == 2:
|
||||
cos = cos[None, None, :, :]
|
||||
sin = sin[None, None, :, :]
|
||||
elif sequence_dim == 1:
|
||||
cos = cos[None, :, None, :]
|
||||
sin = sin[None, :, None, :]
|
||||
else:
|
||||
raise ValueError(f"sequence_dim must be 1 or 2, got {sequence_dim}")
|
||||
|
||||
if use_real_unbind_dim == -1:
|
||||
# Used for flux, cogvideox, hunyuan-dit
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -26,6 +26,7 @@ from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
|
||||
@@ -106,7 +107,9 @@ class WanSelfAttention(nn.Module):
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
parallel_attention=False) -> None:
|
||||
parallel_attention=False,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "") -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -118,10 +121,10 @@ class WanSelfAttention(nn.Module):
|
||||
self.parallel_attention = parallel_attention
|
||||
|
||||
# layers
|
||||
self.to_q = ReplicatedLinear(dim, dim)
|
||||
self.to_k = ReplicatedLinear(dim, dim)
|
||||
self.to_v = ReplicatedLinear(dim, dim)
|
||||
self.to_out = ReplicatedLinear(dim, dim)
|
||||
self.to_q = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_q")
|
||||
self.to_k = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_k")
|
||||
self.to_v = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_v")
|
||||
self.to_out = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_out")
|
||||
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
@@ -194,13 +197,15 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None
|
||||
| None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
supported_attention_backends, quant_config=quant_config, prefix=prefix)
|
||||
|
||||
self.add_k_proj = ReplicatedLinear(dim, dim)
|
||||
self.add_v_proj = ReplicatedLinear(dim, dim)
|
||||
self.add_k_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_k_proj")
|
||||
self.add_v_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_v_proj")
|
||||
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
@@ -246,16 +251,17 @@ class WanTransformerBlock(nn.Module):
|
||||
added_kv_proj_dim: int | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_q")
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_k")
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_v")
|
||||
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_out")
|
||||
self.attn1 = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
@@ -290,13 +296,17 @@ class WanTransformerBlock(nn.Module):
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn2")
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn2")
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
@@ -306,7 +316,7 @@ class WanTransformerBlock(nn.Module):
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
|
||||
self.mlp_residual = ScaleResidual()
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
@@ -406,17 +416,17 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
added_kv_proj_dim: int | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_gate_compress = ReplicatedLinear(dim, dim, bias=True)
|
||||
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_q")
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_k")
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_v")
|
||||
self.to_gate_compress = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_gate_compress")
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_out")
|
||||
self.attn1 = DistributedAttention_VSA(
|
||||
num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
@@ -451,13 +461,17 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn2")
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn2")
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
@@ -467,7 +481,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
|
||||
self.mlp_residual = ScaleResidual()
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
@@ -556,6 +570,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(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
|
||||
@@ -594,6 +609,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
config.eps,
|
||||
config.added_kv_proj_dim,
|
||||
self._supported_attention_backends,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{config.prefix}.blocks.{i}")
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""HF-backed Mistral3 text encoder wrapper for full Flux2."""
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
|
||||
|
||||
class Mistral3ForConditionalGeneration(TextEncoder):
|
||||
"""Loads the Transformers Mistral3 implementation for Flux2 text encoding."""
|
||||
|
||||
supports_hf_from_pretrained = True
|
||||
|
||||
def __init__(self, config: Mistral3TextConfig) -> None:
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
@classmethod
|
||||
def from_pretrained_local(
|
||||
cls,
|
||||
model_path: str,
|
||||
model_config: Mistral3TextConfig,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> nn.Module:
|
||||
from transformers import AutoModelForImageTextToText
|
||||
|
||||
model = AutoModelForImageTextToText.from_pretrained(
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval()
|
||||
if device.type != "cpu":
|
||||
model = model.to(device)
|
||||
return model
|
||||
|
||||
def forward(self, *args: Any, **kwargs: Any) -> Any:
|
||||
raise NotImplementedError(
|
||||
"Mistral3ForConditionalGeneration is loaded through Transformers "
|
||||
"via from_pretrained_local()."
|
||||
)
|
||||
|
||||
|
||||
EntryClass = Mistral3ForConditionalGeneration
|
||||
@@ -0,0 +1,461 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Ported from SGLang: python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py
|
||||
"""Qwen3 causal LM text encoder for FastVideo diffusion models (e.g. Flux2 Klein)."""
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.distributed import get_tp_world_size
|
||||
from fastvideo.layers.activation import SiluAndMul
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.layers.rotary_embedding import get_rope
|
||||
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.loader.weight_utils import (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
|
||||
|
||||
class Qwen3MLP(nn.Module):
|
||||
"""Qwen3 MLP with SwiGLU activation and tensor parallelism."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
intermediate_size: int,
|
||||
hidden_act: str,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
bias: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_sizes=[intermediate_size] * 2,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.gate_up_proj",
|
||||
)
|
||||
self.down_proj = RowParallelLinear(
|
||||
input_size=intermediate_size,
|
||||
output_size=hidden_size,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.down_proj",
|
||||
)
|
||||
if hidden_act != "silu":
|
||||
raise ValueError(
|
||||
f"Unsupported activation: {hidden_act}. Only silu is supported."
|
||||
)
|
||||
self.act_fn = SiluAndMul()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x, _ = self.gate_up_proj(x)
|
||||
x = self.act_fn(x)
|
||||
x, _ = self.down_proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class Qwen3Attention(nn.Module):
|
||||
"""Qwen3 attention with QK-Norm and tensor parallelism.
|
||||
|
||||
Key difference from LLaMA: RMSNorm is applied to Q and K before attention.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3TextConfig,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
rope_theta: float = 1000000.0,
|
||||
rope_scaling: dict[str, Any] | None = None,
|
||||
max_position_embeddings: int = 40960,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
bias: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
tp_size = get_tp_world_size()
|
||||
self.total_num_heads = num_heads
|
||||
assert self.total_num_heads % tp_size == 0
|
||||
self.num_heads = self.total_num_heads // tp_size
|
||||
self.total_num_kv_heads = num_kv_heads
|
||||
if self.total_num_kv_heads >= tp_size:
|
||||
assert self.total_num_kv_heads % tp_size == 0
|
||||
else:
|
||||
assert tp_size % self.total_num_kv_heads == 0
|
||||
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
|
||||
|
||||
self.head_dim = getattr(
|
||||
config, "head_dim", self.hidden_size // self.total_num_heads
|
||||
)
|
||||
self.rotary_dim = self.head_dim
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.rope_theta = rope_theta
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=self.total_num_heads,
|
||||
total_num_kv_heads=self.total_num_kv_heads,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
self.o_proj = RowParallelLinear(
|
||||
input_size=self.total_num_heads * self.head_dim,
|
||||
output_size=hidden_size,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
rms_norm_eps = getattr(config, "rms_norm_eps", 1e-6)
|
||||
self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps)
|
||||
|
||||
self.rotary_emb = get_rope(
|
||||
self.head_dim,
|
||||
rotary_dim=self.rotary_dim,
|
||||
max_position=max_position_embeddings,
|
||||
base=int(rope_theta),
|
||||
rope_scaling=rope_scaling,
|
||||
is_neox_style=True,
|
||||
)
|
||||
|
||||
self.attn = LocalAttention(
|
||||
self.num_heads,
|
||||
self.head_dim,
|
||||
self.num_kv_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=True,
|
||||
supported_attention_backends=config._supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
|
||||
batch_size, seq_len = q.shape[0], q.shape[1]
|
||||
q = q.reshape(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
q = q.reshape(batch_size, seq_len, -1)
|
||||
k = k.reshape(batch_size, seq_len, -1)
|
||||
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
|
||||
q = q.reshape(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
|
||||
if attention_mask is None:
|
||||
attn_output = self.attn(q, k, v)
|
||||
else:
|
||||
q_sdpa = q.transpose(1, 2)
|
||||
k_sdpa = k.transpose(1, 2)
|
||||
v_sdpa = v.transpose(1, 2)
|
||||
causal_mask = torch.ones(
|
||||
seq_len,
|
||||
seq_len,
|
||||
device=q.device,
|
||||
dtype=torch.bool,
|
||||
).tril()
|
||||
key_mask = attention_mask.to(device=q.device, dtype=torch.bool)
|
||||
attn_mask = causal_mask[None, None, :, :] & key_mask[:, None, None, :]
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
q_sdpa,
|
||||
k_sdpa,
|
||||
v_sdpa,
|
||||
attn_mask=attn_mask,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
enable_gqa=self.num_heads != self.num_kv_heads,
|
||||
).transpose(1, 2)
|
||||
|
||||
attn_output = attn_output.reshape(batch_size, seq_len, -1)
|
||||
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
|
||||
|
||||
class Qwen3DecoderLayer(nn.Module):
|
||||
"""Qwen3 transformer decoder layer."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3TextConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = config.hidden_size
|
||||
rope_theta = getattr(config, "rope_theta", 1000000.0)
|
||||
rope_scaling = getattr(config, "rope_scaling", None)
|
||||
max_position_embeddings = getattr(config, "max_position_embeddings", 40960)
|
||||
attention_bias = getattr(config, "attention_bias", False)
|
||||
|
||||
self.self_attn = Qwen3Attention(
|
||||
config=config,
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=config.num_attention_heads,
|
||||
num_kv_heads=getattr(
|
||||
config, "num_key_value_heads", config.num_attention_heads
|
||||
),
|
||||
rope_theta=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
max_position_embeddings=max_position_embeddings,
|
||||
quant_config=quant_config,
|
||||
bias=attention_bias,
|
||||
prefix=f"{prefix}.self_attn",
|
||||
)
|
||||
self.mlp = Qwen3MLP(
|
||||
hidden_size=self.hidden_size,
|
||||
intermediate_size=config.intermediate_size,
|
||||
hidden_act=config.hidden_act,
|
||||
quant_config=quant_config,
|
||||
bias=getattr(config, "mlp_bias", False),
|
||||
prefix=f"{prefix}.mlp",
|
||||
)
|
||||
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = RMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
residual: torch.Tensor | None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
else:
|
||||
hidden_states, residual = self.input_layernorm(hidden_states, residual)
|
||||
|
||||
hidden_states = self.self_attn(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class Qwen3ForCausalLM(TextEncoder):
|
||||
"""Qwen3 causal language model for text encoding in diffusion models (e.g. Flux2 Klein).
|
||||
|
||||
Features:
|
||||
- Tensor parallelism support
|
||||
- FlashAttention/SDPA support via LocalAttention
|
||||
- QK-Norm for better training stability
|
||||
- output_hidden_states for Klein (layers 9, 18, 27)
|
||||
"""
|
||||
|
||||
supports_hf_from_pretrained = True
|
||||
|
||||
def __init__(self, config: Qwen3TextConfig) -> None:
|
||||
super().__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.quant_config = getattr(config, "quant_config", None)
|
||||
|
||||
if getattr(config, "lora_config", None) is not None:
|
||||
max_loras = getattr(config.lora_config, "max_loras", 1)
|
||||
lora_vocab_size = getattr(config.lora_config, "lora_extra_vocab_size", 1)
|
||||
lora_vocab = lora_vocab_size * max_loras
|
||||
else:
|
||||
lora_vocab = 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.org_vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
self.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=self.quant_config,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
Qwen3DecoderLayer(
|
||||
config=config,
|
||||
quant_config=self.quant_config,
|
||||
prefix=f"{config.prefix}.layers.{i}",
|
||||
)
|
||||
for i in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained_local(
|
||||
cls,
|
||||
model_path: str,
|
||||
model_config: Qwen3TextConfig,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> nn.Module:
|
||||
from transformers import AutoModelForCausalLM
|
||||
|
||||
if device.type == "cpu" and torch.cuda.is_available():
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
device = get_local_torch_device()
|
||||
|
||||
return AutoModelForCausalLM.from_pretrained(
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval().to(device)
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.embed_tokens(input_ids)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None = None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs: Any,
|
||||
) -> BaseEncoderOutput:
|
||||
output_hidden_states = (
|
||||
output_hidden_states
|
||||
if output_hidden_states is not None
|
||||
else self.config.output_hidden_states
|
||||
)
|
||||
|
||||
if inputs_embeds is not None:
|
||||
hidden_states = inputs_embeds
|
||||
else:
|
||||
assert input_ids is not None
|
||||
hidden_states = self.get_input_embeddings(input_ids)
|
||||
|
||||
residual = None
|
||||
|
||||
if position_ids is None:
|
||||
position_ids = torch.arange(
|
||||
0, hidden_states.shape[1], device=hidden_states.device
|
||||
).unsqueeze(0)
|
||||
|
||||
all_hidden_states: tuple[Any, ...] | None = (
|
||||
() if output_hidden_states else None
|
||||
)
|
||||
|
||||
for layer in self.layers:
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (
|
||||
(hidden_states,)
|
||||
if residual is None
|
||||
else (hidden_states + residual,)
|
||||
)
|
||||
hidden_states, residual = layer(
|
||||
position_ids,
|
||||
hidden_states,
|
||||
residual,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
hidden_states=all_hidden_states,
|
||||
)
|
||||
|
||||
def load_weights(
|
||||
self, weights: Iterable[tuple[str, torch.Tensor]]
|
||||
) -> set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("model."):
|
||||
name = name[6:]
|
||||
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name:
|
||||
continue
|
||||
|
||||
if "scale" in name:
|
||||
kv_scale_name: str | None = maybe_remap_kv_scale_name(
|
||||
name, params_dict
|
||||
)
|
||||
if kv_scale_name is None:
|
||||
continue
|
||||
name = kv_scale_name
|
||||
|
||||
for (
|
||||
param_name,
|
||||
weight_name,
|
||||
shard_id,
|
||||
) in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
|
||||
loaded_params.add(name)
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
EntryClass = Qwen3ForCausalLM
|
||||
@@ -13,7 +13,7 @@ from typing import cast
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from safetensors.torch import load_file as safetensors_load_file, safe_open
|
||||
from torch.distributed import init_device_mesh
|
||||
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
@@ -573,13 +573,53 @@ class TokenizerLoader(ComponentLoader):
|
||||
# If parsing fails, fall through to AutoTokenizer below.
|
||||
pass
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
resolved_model_path, # "<path to model>/tokenizer"
|
||||
# in v0, this was same string as encoder_name "ClipTextModel"
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
local_files_only=os.path.isdir(resolved_model_path),
|
||||
)
|
||||
# Only Flux2 full's Mistral3 (require_processor=True) must load via
|
||||
# AutoProcessor. Gate the processor_config.json shortcut on that flag so
|
||||
# existing encoders (e.g. HunyuanVideo 1.5 / Qwen2.5-VL) stay on the
|
||||
# historical AutoTokenizer path below even if their tokenizer dir happens
|
||||
# to ship a processor_config.json.
|
||||
require_processor = False
|
||||
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
|
||||
try:
|
||||
require_processor = any(
|
||||
getattr(getattr(cfg, "arch_config", None), "require_processor", False)
|
||||
for cfg in fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
except Exception:
|
||||
require_processor = False
|
||||
|
||||
if require_processor and os.path.exists(os.path.join(resolved_model_path, "processor_config.json")):
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
resolved_model_path,
|
||||
local_files_only=os.path.isdir(resolved_model_path),
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
)
|
||||
logger.info(
|
||||
"Loaded tokenizer/processor from %s: %s",
|
||||
resolved_model_path,
|
||||
processor.__class__.__name__,
|
||||
)
|
||||
return processor
|
||||
|
||||
try:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
resolved_model_path, # "<path to model>/tokenizer"
|
||||
# in v0, this was same string as encoder_name "ClipTextModel"
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
local_files_only=os.path.isdir(resolved_model_path),
|
||||
)
|
||||
except (OSError, ValueError):
|
||||
tokenizer = AutoProcessor.from_pretrained(
|
||||
resolved_model_path,
|
||||
local_files_only=os.path.isdir(resolved_model_path),
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
)
|
||||
logger.info(
|
||||
"Loaded tokenizer/processor from %s: %s",
|
||||
resolved_model_path,
|
||||
tokenizer.__class__.__name__,
|
||||
)
|
||||
return tokenizer
|
||||
padding_side = None
|
||||
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
|
||||
try:
|
||||
@@ -864,6 +904,18 @@ class VocoderLoader(ComponentLoader):
|
||||
return vocoder.eval()
|
||||
|
||||
|
||||
def _collect_safetensors_keys(safetensors_list: list) -> set:
|
||||
"""Collect all weight keys from safetensors files."""
|
||||
all_keys: set[str] = set()
|
||||
for path in safetensors_list:
|
||||
try:
|
||||
with safe_open(path, framework="pt") as f:
|
||||
all_keys.update(f.keys())
|
||||
except Exception as e:
|
||||
logger.warning("Could not read keys from %s: %s", path, e)
|
||||
return all_keys
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
@@ -899,6 +951,12 @@ class TransformerLoader(ComponentLoader):
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
# arch_config can infer architecture from weight keys (e.g. Flux2 layer counts)
|
||||
update_fn = getattr(dit_config.arch_config, "update_from_weight_keys", None)
|
||||
if callable(update_fn):
|
||||
weight_keys = _collect_safetensors_keys(safetensors_list)
|
||||
update_fn(weight_keys)
|
||||
|
||||
# Check if we should use custom initialization weights
|
||||
custom_weights_path = getattr(
|
||||
fastvideo_args, "init_weights_from_safetensors", None
|
||||
|
||||
@@ -332,6 +332,7 @@ def load_model_from_full_model_state_dict(
|
||||
NotImplementedError: If got FSDP with more than 1D.
|
||||
"""
|
||||
meta_sd = model.state_dict()
|
||||
named_parameters = dict(model.named_parameters())
|
||||
sharded_sd = {}
|
||||
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
||||
full_sd_iterator, param_names_mapping) # type: ignore
|
||||
@@ -363,8 +364,25 @@ def load_model_from_full_model_state_dict(
|
||||
)
|
||||
if not hasattr(meta_sharded_param, "device_mesh"):
|
||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
|
||||
sharded_tensor = full_tensor
|
||||
target_param = named_parameters.get(target_param_name)
|
||||
weight_loader = getattr(target_param, "weight_loader", None)
|
||||
# Gated on a shape mismatch: only fused/stacked params with a custom
|
||||
# weight_loader (e.g. Qwen3's merged QKV/gate-up) take this path.
|
||||
# Existing models whose unsharded params match the checkpoint shape
|
||||
# fall through to the original `sharded_tensor = full_tensor` below.
|
||||
if target_param is not None and callable(weight_loader) and tuple(target_param.shape) != tuple(
|
||||
full_tensor.shape):
|
||||
loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape),
|
||||
device=device,
|
||||
dtype=param_dtype),
|
||||
requires_grad=False)
|
||||
for attr_name, attr_value in vars(target_param).items():
|
||||
setattr(loaded_param, attr_name, attr_value)
|
||||
weight_loader(loaded_param, full_tensor)
|
||||
sharded_tensor = loaded_param.data
|
||||
else:
|
||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors.
|
||||
sharded_tensor = full_tensor
|
||||
else:
|
||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||
sharded_tensor = distribute_tensor(
|
||||
|
||||
@@ -42,6 +42,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
|
||||
"Gen3CTransformer3DModel": ("dits", "gen3c", "Gen3CTransformer3DModel"),
|
||||
"Kandinsky5Transformer3DModel": ("dits", "kandinsky5", "Kandinsky5Transformer3DModel"),
|
||||
"Flux2Transformer2DModel": ("dits", "flux_2", "Flux2Transformer2DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -69,6 +70,9 @@ _TEXT_ENCODER_MODELS = {
|
||||
"Qwen2_5_VLForConditionalGeneration":
|
||||
("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
|
||||
"Qwen3ForCausalLM": ("encoders", "qwen3", "Qwen3ForCausalLM"),
|
||||
"Mistral3ForConditionalGeneration":
|
||||
("encoders", "mistral3", "Mistral3ForConditionalGeneration"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
@@ -90,6 +94,7 @@ _VAE_MODELS = {
|
||||
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
|
||||
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
|
||||
"AutoencoderKLFlux2": ("vaes", "flux2vae", "AutoencoderKLFlux2"),
|
||||
# `stable-audio-open-1.0/vae/config.json` ships `_class_name="AutoencoderOobleck"`
|
||||
# (Diffusers' name); FastVideo's class is `OobleckVAE`.
|
||||
"AutoencoderOobleck": ("vaes", "oobleck", "OobleckVAE"),
|
||||
|
||||
@@ -0,0 +1,532 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copyright 2025 The HuggingFace Team. All rights reserved.
|
||||
# Adapted from: huggingface/diffusers `Encoder`/`Decoder` VAE components
|
||||
# at the installed 0.36.0 source surface used by Flux2 Klein.
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.vaes.common import DiagonalGaussianDistribution
|
||||
|
||||
|
||||
@dataclass
|
||||
class AutoencoderKLOutput:
|
||||
latent_dist: DiagonalGaussianDistribution
|
||||
|
||||
def __getitem__(self, idx: int):
|
||||
return (self.latent_dist,)[idx]
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
# Existing local Flux2 parity tests used `vae.encode(x).mean` while
|
||||
# diffusers-style callers use `vae.encode(x).latent_dist.mean`.
|
||||
return getattr(self.latent_dist, name)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecoderOutput:
|
||||
sample: torch.Tensor
|
||||
commit_loss: Optional[torch.Tensor] = None
|
||||
|
||||
def __getitem__(self, idx: int):
|
||||
return (self.sample, self.commit_loss)[idx]
|
||||
|
||||
|
||||
def get_activation(act_fn: str) -> nn.Module:
|
||||
if act_fn in ("swish", "silu"):
|
||||
return nn.SiLU()
|
||||
if act_fn == "mish":
|
||||
return nn.Mish()
|
||||
if act_fn == "gelu":
|
||||
return nn.GELU()
|
||||
if act_fn == "relu":
|
||||
return nn.ReLU()
|
||||
raise ValueError(f"Unsupported activation function: {act_fn}")
|
||||
|
||||
|
||||
class AttnProcessor:
|
||||
def __call__(self, attn: "Attention", hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
return attn._forward(hidden_states, temb=temb)
|
||||
|
||||
|
||||
class AttnAddedKVProcessor(AttnProcessor):
|
||||
pass
|
||||
|
||||
|
||||
ADDED_KV_ATTENTION_PROCESSORS = frozenset({AttnAddedKVProcessor})
|
||||
CROSS_ATTENTION_PROCESSORS = frozenset({AttnProcessor})
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
dropout: float = 0.0,
|
||||
bias: bool = False,
|
||||
upcast_softmax: bool = False,
|
||||
norm_num_groups: Optional[int] = None,
|
||||
spatial_norm_dim: Optional[int] = None,
|
||||
out_bias: bool = True,
|
||||
eps: float = 1e-5,
|
||||
rescale_output_factor: float = 1.0,
|
||||
residual_connection: bool = False,
|
||||
_from_deprecated_attn_block: bool = False,
|
||||
**_: object,
|
||||
):
|
||||
super().__init__()
|
||||
self.inner_dim = dim_head * heads
|
||||
self.query_dim = query_dim
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.scale = dim_head**-0.5
|
||||
self.upcast_softmax = upcast_softmax
|
||||
self.rescale_output_factor = rescale_output_factor
|
||||
self.residual_connection = residual_connection
|
||||
self._from_deprecated_attn_block = _from_deprecated_attn_block
|
||||
self.spatial_norm = None
|
||||
if spatial_norm_dim is not None:
|
||||
raise ValueError("Flux2 VAE does not use spatial attention norm in this port")
|
||||
self.group_norm = (
|
||||
nn.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True)
|
||||
if norm_num_groups is not None
|
||||
else None
|
||||
)
|
||||
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, query_dim, bias=out_bias), nn.Dropout(dropout)])
|
||||
self.processor = AttnProcessor()
|
||||
|
||||
def set_processor(self, processor: AttnProcessor) -> None:
|
||||
self.processor = processor
|
||||
|
||||
def get_processor(self) -> AttnProcessor:
|
||||
return self.processor
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
return self.processor(self, hidden_states, temb=temb)
|
||||
|
||||
def _forward(self, hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states)
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
query = query.view(batch_size, -1, self.heads, self.dim_head).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, self.heads, self.dim_head).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, self.heads, self.dim_head).transpose(1, 2)
|
||||
|
||||
if self.upcast_softmax:
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
value = value.float()
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, scale=self.scale)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, height * width, self.inner_dim)
|
||||
hidden_states = hidden_states.to(self.to_out[0].weight.dtype)
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, channel, height, width)
|
||||
if self.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
return hidden_states / self.rescale_output_factor
|
||||
|
||||
|
||||
class Downsample2D(nn.Module):
|
||||
def __init__(self, channels: int, use_conv: bool = False, out_channels: Optional[int] = None, padding: int = 1, name: str = "conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.padding = padding
|
||||
self.name = name
|
||||
if use_conv:
|
||||
conv = nn.Conv2d(self.channels, self.out_channels, kernel_size=3, stride=2, padding=padding)
|
||||
else:
|
||||
assert self.channels == self.out_channels
|
||||
conv = nn.AvgPool2d(kernel_size=2, stride=2)
|
||||
if name == "conv":
|
||||
self.Conv2d_0 = conv
|
||||
self.conv = conv
|
||||
elif name == "Conv2d_0":
|
||||
self.conv = conv
|
||||
else:
|
||||
self.conv = conv
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
if self.use_conv and self.padding == 0:
|
||||
hidden_states = F.pad(hidden_states, (0, 1, 0, 1), mode="constant", value=0)
|
||||
return self.conv(hidden_states)
|
||||
|
||||
|
||||
class Upsample2D(nn.Module):
|
||||
def __init__(self, channels: int, use_conv: bool = False, out_channels: Optional[int] = None, name: str = "conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_conv_transpose = False
|
||||
self.name = name
|
||||
self.interpolate = True
|
||||
conv = nn.Conv2d(self.channels, self.out_channels, kernel_size=3, padding=1) if use_conv else None
|
||||
if name == "conv":
|
||||
self.conv = conv
|
||||
else:
|
||||
self.Conv2d_0 = conv
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, output_size: Optional[int] = None, *args, **kwargs) -> torch.Tensor:
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
dtype = hidden_states.dtype
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.float()
|
||||
if output_size is None:
|
||||
hidden_states = F.interpolate(hidden_states, scale_factor=2.0, mode="nearest")
|
||||
else:
|
||||
hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest")
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(dtype)
|
||||
if self.use_conv:
|
||||
hidden_states = self.conv(hidden_states) if self.name == "conv" else self.Conv2d_0(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ResnetBlock2D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
dropout: float = 0.0,
|
||||
temb_channels: Optional[int] = 512,
|
||||
groups: int = 32,
|
||||
groups_out: Optional[int] = None,
|
||||
eps: float = 1e-6,
|
||||
non_linearity: str = "swish",
|
||||
time_embedding_norm: str = "default",
|
||||
output_scale_factor: float = 1.0,
|
||||
use_in_shortcut: Optional[bool] = None,
|
||||
conv_shortcut_bias: bool = True,
|
||||
conv_2d_out_channels: Optional[int] = None,
|
||||
**_: object,
|
||||
):
|
||||
super().__init__()
|
||||
if time_embedding_norm not in ("default", "scale_shift"):
|
||||
raise ValueError(f"unknown time_embedding_norm: {time_embedding_norm}")
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.output_scale_factor = output_scale_factor
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
self.norm1 = nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
if temb_channels is not None:
|
||||
self.time_emb_proj = nn.Linear(temb_channels, out_channels if time_embedding_norm == "default" else 2 * out_channels)
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
self.norm2 = nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
conv_2d_out_channels = conv_2d_out_channels or out_channels
|
||||
self.conv2 = nn.Conv2d(out_channels, conv_2d_out_channels, kernel_size=3, stride=1, padding=1)
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
self.use_in_shortcut = self.in_channels != conv_2d_out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
self.conv_shortcut = nn.Conv2d(in_channels, conv_2d_out_channels, kernel_size=1, stride=1, padding=0, bias=conv_shortcut_bias)
|
||||
|
||||
def forward(self, input_tensor: torch.Tensor, temb: Optional[torch.Tensor] = None, *args, **kwargs) -> torch.Tensor:
|
||||
hidden_states = self.norm1(input_tensor)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
if self.time_emb_proj is not None and temb is not None:
|
||||
temb = self.nonlinearity(temb)
|
||||
temb = self.time_emb_proj(temb)[:, :, None, None]
|
||||
if self.time_embedding_norm == "default":
|
||||
if temb is not None:
|
||||
hidden_states = hidden_states + temb
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
else:
|
||||
if temb is None:
|
||||
raise ValueError("temb cannot be None for scale_shift")
|
||||
time_scale, time_shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = hidden_states * (1 + time_scale) + time_shift
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor.contiguous() if self.training else input_tensor)
|
||||
return (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
|
||||
class UNetMidBlock2D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
temb_channels: Optional[int],
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
attn_groups: Optional[int] = None,
|
||||
resnet_pre_norm: bool = True,
|
||||
add_attention: bool = True,
|
||||
attention_head_dim: int = 1,
|
||||
output_scale_factor: float = 1.0,
|
||||
):
|
||||
super().__init__()
|
||||
if resnet_time_scale_shift == "spatial":
|
||||
raise ValueError("Flux2 VAE does not use spatial resnet conditioning in this port")
|
||||
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
self.add_attention = add_attention
|
||||
if attn_groups is None:
|
||||
attn_groups = resnet_groups
|
||||
resnets = [
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
]
|
||||
attentions = []
|
||||
if attention_head_dim is None:
|
||||
attention_head_dim = in_channels
|
||||
for _ in range(num_layers):
|
||||
attentions.append(
|
||||
Attention(
|
||||
in_channels,
|
||||
heads=in_channels // attention_head_dim,
|
||||
dim_head=attention_head_dim,
|
||||
rescale_output_factor=output_scale_factor,
|
||||
eps=resnet_eps,
|
||||
norm_num_groups=attn_groups,
|
||||
residual_connection=True,
|
||||
bias=True,
|
||||
upcast_softmax=True,
|
||||
_from_deprecated_attn_block=True,
|
||||
)
|
||||
if self.add_attention
|
||||
else None
|
||||
)
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
hidden_states = self.resnets[0](hidden_states, temb)
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
hidden_states = attn(hidden_states, temb=temb)
|
||||
hidden_states = resnet(hidden_states, temb)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DownEncoderBlock2D(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, dropout: float = 0.0, num_layers: int = 1, resnet_eps: float = 1e-6, resnet_act_fn: str = "swish", resnet_groups: int = 32, add_downsample: bool = True, downsample_padding: int = 1, **_: object):
|
||||
super().__init__()
|
||||
self.resnets = nn.ModuleList([
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels if i == 0 else out_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=None,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
self.downsamplers = nn.ModuleList([Downsample2D(out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op")]) if add_downsample else None
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=None)
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AttnDownEncoderBlock2D(DownEncoderBlock2D):
|
||||
def __init__(self, in_channels: int, out_channels: int, attention_head_dim: int = 1, **kwargs: object):
|
||||
super().__init__(in_channels=in_channels, out_channels=out_channels, **kwargs)
|
||||
resnet_groups = int(kwargs.get("resnet_groups", 32))
|
||||
resnet_eps = float(kwargs.get("resnet_eps", 1e-6))
|
||||
output_scale_factor = float(kwargs.get("output_scale_factor", 1.0))
|
||||
if attention_head_dim is None:
|
||||
attention_head_dim = out_channels
|
||||
self.attentions = nn.ModuleList([
|
||||
Attention(out_channels, heads=out_channels // attention_head_dim, dim_head=attention_head_dim, rescale_output_factor=output_scale_factor, eps=resnet_eps, norm_num_groups=resnet_groups, residual_connection=True, bias=True, upcast_softmax=True, _from_deprecated_attn_block=True)
|
||||
for _ in self.resnets
|
||||
])
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
for resnet, attn in zip(self.resnets, self.attentions):
|
||||
hidden_states = resnet(hidden_states, temb=None)
|
||||
hidden_states = attn(hidden_states)
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class UpDecoderBlock2D(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, dropout: float = 0.0, num_layers: int = 1, resnet_eps: float = 1e-6, resnet_act_fn: str = "swish", resnet_groups: int = 32, add_upsample: bool = True, temb_channels: Optional[int] = None, **_: object):
|
||||
super().__init__()
|
||||
self.resnets = nn.ModuleList([
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels if i == 0 else out_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
non_linearity=resnet_act_fn,
|
||||
)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) if add_upsample else None
|
||||
self.resolution_idx = None
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=temb)
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = upsampler(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AttnUpDecoderBlock2D(UpDecoderBlock2D):
|
||||
def __init__(self, in_channels: int, out_channels: int, attention_head_dim: int = 1, **kwargs: object):
|
||||
super().__init__(in_channels=in_channels, out_channels=out_channels, **kwargs)
|
||||
resnet_groups = int(kwargs.get("resnet_groups", 32))
|
||||
resnet_eps = float(kwargs.get("resnet_eps", 1e-6))
|
||||
output_scale_factor = float(kwargs.get("output_scale_factor", 1.0))
|
||||
if attention_head_dim is None:
|
||||
attention_head_dim = out_channels
|
||||
self.attentions = nn.ModuleList([
|
||||
Attention(out_channels, heads=out_channels // attention_head_dim, dim_head=attention_head_dim, rescale_output_factor=output_scale_factor, eps=resnet_eps, norm_num_groups=resnet_groups, residual_connection=True, bias=True, upcast_softmax=True, _from_deprecated_attn_block=True)
|
||||
for _ in self.resnets
|
||||
])
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, temb: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
for resnet, attn in zip(self.resnets, self.attentions):
|
||||
hidden_states = resnet(hidden_states, temb=temb)
|
||||
hidden_states = attn(hidden_states, temb=temb)
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = upsampler(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
def get_down_block(down_block_type: str, **kwargs: object) -> nn.Module:
|
||||
if down_block_type == "DownEncoderBlock2D":
|
||||
return DownEncoderBlock2D(**kwargs)
|
||||
if down_block_type == "AttnDownEncoderBlock2D":
|
||||
return AttnDownEncoderBlock2D(**kwargs)
|
||||
raise ValueError(f"Unsupported Flux2 VAE down block type: {down_block_type}")
|
||||
|
||||
|
||||
def get_up_block(up_block_type: str, **kwargs: object) -> nn.Module:
|
||||
kwargs.pop("prev_output_channel", None)
|
||||
kwargs.pop("resolution_idx", None)
|
||||
if up_block_type == "UpDecoderBlock2D":
|
||||
return UpDecoderBlock2D(**kwargs)
|
||||
if up_block_type == "AttnUpDecoderBlock2D":
|
||||
return AttnUpDecoderBlock2D(**kwargs)
|
||||
raise ValueError(f"Unsupported Flux2 VAE up block type: {up_block_type}")
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, in_channels: int = 3, out_channels: int = 3, down_block_types: Tuple[str, ...] = ("DownEncoderBlock2D",), block_out_channels: Tuple[int, ...] = (64,), layers_per_block: int = 2, norm_num_groups: int = 32, act_fn: str = "silu", double_z: bool = True, mid_block_add_attention: bool = True):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
self.conv_in = nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
self.down_blocks.append(get_down_block(down_block_type, num_layers=self.layers_per_block, in_channels=input_channel, out_channels=output_channel, add_downsample=not is_final_block, resnet_eps=1e-6, downsample_padding=0, resnet_act_fn=act_fn, resnet_groups=norm_num_groups, attention_head_dim=output_channel, temb_channels=None))
|
||||
self.mid_block = UNetMidBlock2D(in_channels=block_out_channels[-1], resnet_eps=1e-6, resnet_act_fn=act_fn, output_scale_factor=1, resnet_time_scale_shift="default", attention_head_dim=block_out_channels[-1], resnet_groups=norm_num_groups, temb_channels=None, add_attention=mid_block_add_attention)
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
conv_out_channels = 2 * out_channels if double_z else out_channels
|
||||
self.conv_out = nn.Conv2d(block_out_channels[-1], conv_out_channels, 3, padding=1)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
sample = self.conv_in(sample)
|
||||
for down_block in self.down_blocks:
|
||||
sample = down_block(sample)
|
||||
sample = self.mid_block(sample)
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
return self.conv_out(sample)
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, in_channels: int = 3, out_channels: int = 3, up_block_types: Tuple[str, ...] = ("UpDecoderBlock2D",), block_out_channels: Tuple[int, ...] = (64,), layers_per_block: int = 2, norm_num_groups: int = 32, act_fn: str = "silu", norm_type: str = "group", mid_block_add_attention: bool = True):
|
||||
super().__init__()
|
||||
if norm_type != "group":
|
||||
raise ValueError("Flux2 VAE Decoder only supports group norm in this port")
|
||||
self.layers_per_block = layers_per_block
|
||||
self.conv_in = nn.Conv2d(in_channels, block_out_channels[-1], kernel_size=3, stride=1, padding=1)
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
self.mid_block = UNetMidBlock2D(in_channels=block_out_channels[-1], resnet_eps=1e-6, resnet_act_fn=act_fn, output_scale_factor=1, resnet_time_scale_shift="default", attention_head_dim=block_out_channels[-1], resnet_groups=norm_num_groups, temb_channels=None, add_attention=mid_block_add_attention)
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
self.up_blocks.append(get_up_block(up_block_type, num_layers=self.layers_per_block + 1, in_channels=prev_output_channel, out_channels=output_channel, prev_output_channel=prev_output_channel, add_upsample=not is_final_block, resnet_eps=1e-6, resnet_act_fn=act_fn, resnet_groups=norm_num_groups, attention_head_dim=output_channel, temb_channels=None, resnet_time_scale_shift=norm_type))
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, 3, padding=1)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, sample: torch.Tensor, latent_embeds: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
sample = self.conv_in(sample)
|
||||
sample = self.mid_block(sample, latent_embeds)
|
||||
for up_block in self.up_blocks:
|
||||
sample = up_block(sample, latent_embeds)
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
return self.conv_out(sample)
|
||||
@@ -0,0 +1,533 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
import math
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.models.vaes.common import (
|
||||
DiagonalGaussianDistribution,
|
||||
ParallelTiledVAE,
|
||||
)
|
||||
from fastvideo.models.vaes.flux2_components import (
|
||||
ADDED_KV_ATTENTION_PROCESSORS,
|
||||
CROSS_ATTENTION_PROCESSORS,
|
||||
Attention,
|
||||
AttnAddedKVProcessor,
|
||||
AttnProcessor,
|
||||
AutoencoderKLOutput,
|
||||
Decoder,
|
||||
DecoderOutput,
|
||||
Encoder,
|
||||
)
|
||||
|
||||
AttentionProcessor = AttnProcessor
|
||||
|
||||
|
||||
class AutoencoderKLFlux2(nn.Module, ParallelTiledVAE):
|
||||
r"""
|
||||
A VAE model with KL loss for encoding images into latents and decoding latent representations into images.
|
||||
|
||||
This model inherits from [`ParallelTiledVAE`] for tiling support and uses standard diffusers
|
||||
Encoder/Decoder components for Flux2 image generation.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_no_split_modules = ["Attention", "ResnetBlock2D"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Flux2VAEConfig,
|
||||
):
|
||||
nn.Module.__init__(self)
|
||||
ParallelTiledVAE.__init__(self, config=config)
|
||||
|
||||
self.config = config
|
||||
arch_config = config.arch_config
|
||||
|
||||
in_channels: int = arch_config.in_channels
|
||||
out_channels: int = arch_config.out_channels
|
||||
down_block_types: Tuple[str, ...] = arch_config.down_block_types
|
||||
up_block_types: Tuple[str, ...] = arch_config.up_block_types
|
||||
block_out_channels: Tuple[int, ...] = arch_config.block_out_channels
|
||||
layers_per_block: int = arch_config.layers_per_block
|
||||
act_fn: str = arch_config.act_fn
|
||||
latent_channels: int = arch_config.latent_channels
|
||||
norm_num_groups: int = arch_config.norm_num_groups
|
||||
sample_size: int = arch_config.sample_size
|
||||
force_upcast: bool = arch_config.force_upcast
|
||||
use_quant_conv: bool = arch_config.use_quant_conv
|
||||
use_post_quant_conv: bool = arch_config.use_post_quant_conv
|
||||
mid_block_add_attention: bool = arch_config.mid_block_add_attention
|
||||
batch_norm_eps: float = arch_config.batch_norm_eps
|
||||
batch_norm_momentum: float = arch_config.batch_norm_momentum
|
||||
patch_size: Tuple[int, int] = arch_config.patch_size
|
||||
|
||||
# pass init params to Encoder
|
||||
self.encoder = Encoder(
|
||||
in_channels=in_channels,
|
||||
out_channels=latent_channels,
|
||||
down_block_types=down_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
act_fn=act_fn,
|
||||
norm_num_groups=norm_num_groups,
|
||||
double_z=True,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
# pass init params to Decoder
|
||||
self.decoder = Decoder(
|
||||
in_channels=latent_channels,
|
||||
out_channels=out_channels,
|
||||
up_block_types=up_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
self.quant_conv = (
|
||||
nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1)
|
||||
if use_quant_conv
|
||||
else None
|
||||
)
|
||||
self.post_quant_conv = (
|
||||
nn.Conv2d(latent_channels, latent_channels, 1)
|
||||
if use_post_quant_conv
|
||||
else None
|
||||
)
|
||||
|
||||
self.bn = nn.BatchNorm2d(
|
||||
math.prod(patch_size) * latent_channels,
|
||||
eps=batch_norm_eps,
|
||||
momentum=batch_norm_momentum,
|
||||
affine=False,
|
||||
track_running_stats=True,
|
||||
)
|
||||
|
||||
self.use_slicing = False
|
||||
self.use_tiling = False
|
||||
|
||||
# only relevant if vae tiling is enabled
|
||||
self.tile_sample_min_size = sample_size
|
||||
sample_size_val = (
|
||||
sample_size[0]
|
||||
if isinstance(sample_size, (list, tuple))
|
||||
else sample_size
|
||||
)
|
||||
self.tile_latent_min_size = int(
|
||||
sample_size_val / (2 ** (len(block_out_channels) - 1))
|
||||
)
|
||||
self.tile_overlap_factor = 0.25
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(
|
||||
name: str,
|
||||
module: torch.nn.Module,
|
||||
processors: Dict[str, AttentionProcessor],
|
||||
):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor()
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(
|
||||
self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]
|
||||
):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(
|
||||
proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
processor = AttnAddedKVProcessor()
|
||||
elif all(
|
||||
proc.__class__ in CROSS_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
self.set_attn_processor(processor)
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, height, width = x.shape
|
||||
|
||||
if self.use_tiling and (
|
||||
width > self.tile_sample_min_size or height > self.tile_sample_min_size
|
||||
):
|
||||
return self._tiled_encode(x)
|
||||
|
||||
enc = self.encoder(x)
|
||||
if self.quant_conv is not None:
|
||||
enc = self.quant_conv(enc)
|
||||
|
||||
return enc
|
||||
|
||||
def encode(
|
||||
self, x: torch.Tensor, return_dict: bool = True
|
||||
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
|
||||
"""
|
||||
Encode a batch of images into latents.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`): Input batch of images.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
The latent representations of the encoded images. If `return_dict` is True, a
|
||||
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
|
||||
"""
|
||||
|
||||
if x.ndim == 5:
|
||||
assert x.shape[2] == 1
|
||||
x = x.squeeze(2)
|
||||
|
||||
if self.use_slicing and x.shape[0] > 1:
|
||||
encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
|
||||
h = torch.cat(encoded_slices)
|
||||
else:
|
||||
h = self._encode(x)
|
||||
|
||||
posterior = DiagonalGaussianDistribution(h)
|
||||
if not return_dict:
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def _decode(
|
||||
self, z: torch.Tensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.Tensor]:
|
||||
if self.use_tiling and (
|
||||
z.shape[-1] > self.tile_latent_min_size
|
||||
or z.shape[-2] > self.tile_latent_min_size
|
||||
):
|
||||
return self.tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
if self.post_quant_conv is not None:
|
||||
z = self.post_quant_conv(z)
|
||||
|
||||
dec = self.decoder(z)
|
||||
|
||||
if not return_dict:
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True, generator=None
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Decode a batch of images.
|
||||
|
||||
Args:
|
||||
z (`torch.Tensor`): Input batch of latent vectors.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.vae.DecoderOutput`] or `tuple`:
|
||||
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
|
||||
"""
|
||||
if self.use_slicing and z.shape[0] > 1:
|
||||
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
|
||||
decoded = torch.cat(decoded_slices)
|
||||
else:
|
||||
decoded = self._decode(z).sample
|
||||
|
||||
if not return_dict:
|
||||
return (decoded,)
|
||||
|
||||
return DecoderOutput(sample=decoded)
|
||||
|
||||
def blend_v(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[2], b.shape[2], blend_extent)
|
||||
for y in range(blend_extent):
|
||||
b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[
|
||||
:, :, y, :
|
||||
] * (y / blend_extent)
|
||||
return b
|
||||
|
||||
def blend_h(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[3], b.shape[3], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[
|
||||
:, :, :, x
|
||||
] * (x / blend_extent)
|
||||
return b
|
||||
|
||||
def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
r"""Encode a batch of images using a tiled encoder.
|
||||
|
||||
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
|
||||
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
|
||||
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
|
||||
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
|
||||
output, but they should be much less noticeable.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`): Input batch of images.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The latent representation of the encoded videos.
|
||||
"""
|
||||
|
||||
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_latent_min_size - blend_extent
|
||||
|
||||
# Split the image into 512x512 tiles and encode them separately.
|
||||
rows = []
|
||||
for i in range(0, x.shape[2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, x.shape[3], overlap_size):
|
||||
tile = x[
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_sample_min_size,
|
||||
j : j + self.tile_sample_min_size,
|
||||
]
|
||||
tile = self.encoder(tile)
|
||||
if self.quant_conv is not None:
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=3))
|
||||
|
||||
enc = torch.cat(result_rows, dim=2)
|
||||
return enc
|
||||
|
||||
def tiled_encode(
|
||||
self, x: torch.Tensor, return_dict: bool = True
|
||||
) -> AutoencoderKLOutput:
|
||||
r"""Encode a batch of images using a tiled encoder.
|
||||
|
||||
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
|
||||
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
|
||||
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
|
||||
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
|
||||
output, but they should be much less noticeable.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`): Input batch of images.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
|
||||
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
|
||||
`tuple` is returned.
|
||||
"""
|
||||
deprecation_message = (
|
||||
"The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the "
|
||||
"implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able "
|
||||
"to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value."
|
||||
)
|
||||
|
||||
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_latent_min_size - blend_extent
|
||||
|
||||
# Split the image into 512x512 tiles and encode them separately.
|
||||
rows = []
|
||||
for i in range(0, x.shape[2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, x.shape[3], overlap_size):
|
||||
tile = x[
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_sample_min_size,
|
||||
j : j + self.tile_sample_min_size,
|
||||
]
|
||||
tile = self.encoder(tile)
|
||||
if self.quant_conv is not None:
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=3))
|
||||
|
||||
moments = torch.cat(result_rows, dim=2)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def tiled_decode(
|
||||
self, z: torch.Tensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.Tensor]:
|
||||
r"""
|
||||
Decode a batch of images using a tiled decoder.
|
||||
|
||||
Args:
|
||||
z (`torch.Tensor`): Input batch of latent vectors.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.vae.DecoderOutput`] or `tuple`:
|
||||
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
"""
|
||||
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_sample_min_size - blend_extent
|
||||
|
||||
# Split z into overlapping 64x64 tiles and decode them separately.
|
||||
# The tiles have an overlap to avoid seams between tiles.
|
||||
rows = []
|
||||
for i in range(0, z.shape[2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, z.shape[3], overlap_size):
|
||||
tile = z[
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_latent_min_size,
|
||||
j : j + self.tile_latent_min_size,
|
||||
]
|
||||
if self.post_quant_conv is not None:
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
row.append(decoded)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=3))
|
||||
|
||||
dec = torch.cat(result_rows, dim=2)
|
||||
if not return_dict:
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
sample_posterior: bool = False,
|
||||
return_dict: bool = True,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> Union[DecoderOutput, torch.Tensor]:
|
||||
r"""
|
||||
Args:
|
||||
sample (`torch.Tensor`): Input sample.
|
||||
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||
Whether to sample from the posterior.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
|
||||
"""
|
||||
x = sample
|
||||
posterior = self.encode(x).latent_dist
|
||||
if sample_posterior:
|
||||
z = posterior.sample(generator=generator)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z).sample
|
||||
|
||||
if not return_dict:
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
|
||||
EntryClass = AutoencoderKLFlux2
|
||||
@@ -0,0 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Flux2 pipeline module."""
|
||||
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_pipeline import Flux2Pipeline
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_klein_pipeline import Flux2KleinPipeline
|
||||
|
||||
__all__ = ["Flux2Pipeline", "Flux2KleinPipeline"]
|
||||
@@ -0,0 +1,17 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
"""
|
||||
Flux2 Klein image generation pipeline (distilled, 4-step, no guidance).
|
||||
"""
|
||||
|
||||
from fastvideo.configs.pipelines.flux_2 import Flux2KleinPipelineConfig
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_pipeline import Flux2Pipeline
|
||||
|
||||
|
||||
class Flux2KleinPipeline(Flux2Pipeline):
|
||||
"""Flux2 Klein image diffusion pipeline (distilled, 4-step, no guidance)."""
|
||||
|
||||
pipeline_config_cls: type[Flux2KleinPipelineConfig] = Flux2KleinPipelineConfig
|
||||
|
||||
|
||||
EntryClass = Flux2KleinPipeline
|
||||
@@ -0,0 +1,138 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Flux2 latent preparation stage using packed 2x2 layout.
|
||||
|
||||
Flux2 uses packed latents: transformer sees 128 channels (32*4) with half
|
||||
spatial resolution; after denoising we unpatchify to 32 channels and full
|
||||
spatial for VAE decode. This stage prepares (B, 128, T, H//2, W//2).
|
||||
"""
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
|
||||
|
||||
|
||||
class Flux2LatentPreparationStage(LatentPreparationStage):
|
||||
"""
|
||||
Latent preparation for Flux2: packed layout with half spatial dimensions.
|
||||
|
||||
Matches diffusers Flux2Pipeline.prepare_latents: shape is
|
||||
(B, num_channels_latents, T, H_latent//2, W_latent//2) so the transformer
|
||||
sees 128 channels and half spatial; after denoising we unpatchify to
|
||||
(B, 32, H_latent, W_latent) before VAE.
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Prepare latents with Flux2 packed half-spatial shape."""
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
latent_num_frames = None
|
||||
if hasattr(self, "adjust_video_length"):
|
||||
latent_num_frames = self.adjust_video_length(batch, fastvideo_args)
|
||||
|
||||
if not batch.prompt_embeds:
|
||||
if batch.keyboard_cond is not None:
|
||||
batch_size = batch.keyboard_cond.shape[0]
|
||||
elif batch.mouse_cond is not None:
|
||||
batch_size = batch.mouse_cond.shape[0]
|
||||
elif batch.image_embeds:
|
||||
batch_size = batch.image_embeds[0].shape[0]
|
||||
else:
|
||||
batch_size = 1
|
||||
elif isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
if not batch.prompt_embeds:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
device = get_local_torch_device()
|
||||
dummy_prompt = torch.zeros(
|
||||
batch_size,
|
||||
0,
|
||||
self.transformer.hidden_size,
|
||||
device=device,
|
||||
dtype=transformer_dtype,
|
||||
)
|
||||
batch.prompt_embeds = [dummy_prompt]
|
||||
batch.negative_prompt_embeds = []
|
||||
batch.do_classifier_free_guidance = False
|
||||
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
device = get_local_torch_device()
|
||||
generator = batch.generator
|
||||
latents = batch.latents
|
||||
num_frames = (latent_num_frames if latent_num_frames is not None else batch.num_frames)
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
|
||||
scale = vae_arch.spatial_compression_ratio
|
||||
# Flux2 packed: half spatial (2x2 patch packing)
|
||||
latent_h = (height // scale) // 2
|
||||
latent_w = (width // scale) // 2
|
||||
|
||||
if self.use_btchw_layout:
|
||||
shape = (
|
||||
batch_size,
|
||||
num_frames,
|
||||
self.transformer.num_channels_latents,
|
||||
latent_h,
|
||||
latent_w,
|
||||
)
|
||||
bcthw_shape = tuple(shape[i] for i in [0, 2, 1, 3, 4])
|
||||
else:
|
||||
shape = (
|
||||
batch_size,
|
||||
self.transformer.num_channels_latents,
|
||||
num_frames,
|
||||
latent_h,
|
||||
latent_w,
|
||||
)
|
||||
bcthw_shape = shape
|
||||
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(f"You have passed a list of generators of length {len(generator)}, "
|
||||
f"but requested an effective batch size of {batch_size}.")
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
is_longcat_refine = (batch.refine_from is not None or batch.stage1_video is not None)
|
||||
if (not is_longcat_refine) and hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = bcthw_shape
|
||||
latent_ids = torch.cartesian_prod(
|
||||
torch.arange(num_frames, device=device),
|
||||
torch.arange(latent_h, device=device),
|
||||
torch.arange(latent_w, device=device),
|
||||
torch.arange(1, device=device),
|
||||
)
|
||||
batch.extra["flux2_img_ids"] = latent_ids.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
# Flux2 mu depends on image_seq_len; use packed spatial size
|
||||
batch.n_tokens = latent_h * latent_w
|
||||
return batch
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copied and adapted from: https://github.com/sglang-ai/sglang
|
||||
"""
|
||||
Flux2 image generation pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Flux2 image diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_latent_preparation import (
|
||||
Flux2LatentPreparationStage, )
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_timestep_preparation import (
|
||||
Flux2TimestepPreparationStage, )
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_text_encoding import (
|
||||
Flux2TextEncodingStage, )
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Flux2Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Flux2 image diffusion pipeline with LoRA support.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=Flux2TextEncodingStage(
|
||||
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="latent_preparation_stage",
|
||||
stage=Flux2LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=Flux2TimestepPreparationStage(scheduler=self.get_module("scheduler"), ),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
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,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = Flux2Pipeline
|
||||
@@ -0,0 +1,161 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Flux2 text encoding stages."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
|
||||
FLUX2_SYSTEM_MESSAGE = ("You are an AI that reasons about image descriptions. You give structured "
|
||||
"responses focusing on object relationships, object\nattribution and actions "
|
||||
"without speculation.")
|
||||
|
||||
|
||||
def _format_flux2_full_input(prompts: list[str], system_message: str) -> list[list[dict[str, Any]]]:
|
||||
return [[
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": system_message
|
||||
}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": prompt.replace("[IMG]", "")
|
||||
}],
|
||||
},
|
||||
] for prompt in prompts]
|
||||
|
||||
|
||||
def _prepare_flux2_text_ids(prompt_embeds: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, seq_len, _ = prompt_embeds.shape
|
||||
text_ids = torch.cartesian_prod(
|
||||
torch.arange(1, device=prompt_embeds.device),
|
||||
torch.arange(1, device=prompt_embeds.device),
|
||||
torch.arange(1, device=prompt_embeds.device),
|
||||
torch.arange(seq_len, device=prompt_embeds.device),
|
||||
)
|
||||
return text_ids.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
|
||||
|
||||
class Flux2TextEncodingStage(TextEncodingStage):
|
||||
"""Text encoding for Flux2 full and Klein variants."""
|
||||
|
||||
def _uses_embedded_guidance(self, fastvideo_args: FastVideoArgs) -> bool:
|
||||
return getattr(fastvideo_args.pipeline_config, "embedded_cfg_scale", None) is not None
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if self._uses_embedded_guidance(fastvideo_args):
|
||||
batch.do_classifier_free_guidance = False
|
||||
batch.negative_prompt_embeds = []
|
||||
|
||||
if batch.prompt_embeds is not None and len(batch.prompt_embeds) > 0:
|
||||
if "flux2_txt_ids" not in batch.extra:
|
||||
batch.extra["flux2_txt_ids"] = _prepare_flux2_text_ids(batch.prompt_embeds[0])
|
||||
return batch
|
||||
|
||||
if getattr(fastvideo_args.pipeline_config, "flux2_text_encoder_type", "") != "mistral3":
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
assert batch.prompt is not None
|
||||
prompt_embeds, attention_mask = self.encode_flux2_full_text(
|
||||
batch.prompt,
|
||||
fastvideo_args,
|
||||
max_length=batch.max_sequence_length,
|
||||
)
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
batch.extra["flux2_txt_ids"] = _prepare_flux2_text_ids(prompt_embeds)
|
||||
if batch.prompt_attention_mask is not None:
|
||||
batch.prompt_attention_mask.append(attention_mask)
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_flux2_full_text(
|
||||
self,
|
||||
text: str | list[str],
|
||||
fastvideo_args: FastVideoArgs,
|
||||
max_length: int | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
tokenizer = self.tokenizers[0]
|
||||
text_encoder = self.text_encoders[0]
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
arch_config = encoder_config.arch_config
|
||||
|
||||
prompts = [text] if isinstance(text, str) else text
|
||||
max_sequence_length = max_length or getattr(arch_config, "text_len", 512) or 512
|
||||
hidden_state_layers = getattr(
|
||||
fastvideo_args.pipeline_config,
|
||||
"text_encoder_out_layers",
|
||||
(10, 20, 30),
|
||||
)
|
||||
system_message = getattr(
|
||||
fastvideo_args.pipeline_config,
|
||||
"flux2_system_message",
|
||||
FLUX2_SYSTEM_MESSAGE,
|
||||
)
|
||||
|
||||
messages = _format_flux2_full_input(prompts, system_message)
|
||||
inputs = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
add_generation_prompt=False,
|
||||
tokenize=True,
|
||||
return_dict=True,
|
||||
return_tensors="pt",
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=max_sequence_length,
|
||||
)
|
||||
|
||||
try:
|
||||
encoder_device = next(text_encoder.parameters()).device
|
||||
except StopIteration:
|
||||
encoder_device = get_local_torch_device()
|
||||
encoder_dtype = getattr(text_encoder, "dtype", None)
|
||||
|
||||
input_ids = inputs["input_ids"].to(encoder_device)
|
||||
attention_mask = inputs["attention_mask"].to(encoder_device)
|
||||
|
||||
forward_kwargs: dict[str, Any] = {
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"output_hidden_states": True,
|
||||
"use_cache": False,
|
||||
}
|
||||
if "pixel_values" in inputs:
|
||||
forward_kwargs["pixel_values"] = inputs["pixel_values"].to(
|
||||
device=encoder_device,
|
||||
dtype=encoder_dtype or torch.bfloat16,
|
||||
)
|
||||
if "image_sizes" in inputs:
|
||||
forward_kwargs["image_sizes"] = inputs["image_sizes"].to(encoder_device)
|
||||
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = text_encoder(**forward_kwargs)
|
||||
|
||||
if outputs.hidden_states is None:
|
||||
raise ValueError("Full Flux2 requires output_hidden_states=True from text encoder")
|
||||
|
||||
stacked = torch.stack([outputs.hidden_states[k] for k in hidden_state_layers], dim=1)
|
||||
if encoder_dtype is not None:
|
||||
stacked = stacked.to(dtype=encoder_dtype)
|
||||
batch_size, num_layers, seq_len, hidden_dim = stacked.shape
|
||||
prompt_embeds = stacked.permute(0, 2, 1, 3).reshape(
|
||||
batch_size,
|
||||
seq_len,
|
||||
num_layers * hidden_dim,
|
||||
)
|
||||
return prompt_embeds, attention_mask
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Flux2-specific timestep preparation."""
|
||||
|
||||
import inspect
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
|
||||
|
||||
|
||||
def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float:
|
||||
"""
|
||||
Resolution-dependent mu for Flux2 flow-match scheduler.
|
||||
From Black Forest Labs flux2 official repo: sampling.compute_empirical_mu.
|
||||
"""
|
||||
a1, b1 = 8.73809524e-05, 1.89833333
|
||||
a2, b2 = 0.00016927, 0.45666666
|
||||
|
||||
if image_seq_len > 4300:
|
||||
return float(a2 * image_seq_len + b2)
|
||||
|
||||
m_200 = a2 * image_seq_len + b2
|
||||
m_10 = a1 * image_seq_len + b1
|
||||
a = (m_200 - m_10) / 190.0
|
||||
b = m_200 - 200.0 * a
|
||||
return float(a * num_steps + b)
|
||||
|
||||
|
||||
class Flux2TimestepPreparationStage(TimestepPreparationStage):
|
||||
"""Flux2 timestep preparation matching the Diffusers Flux2 schedule."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
scheduler = self.scheduler
|
||||
device = get_local_torch_device()
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
timesteps = batch.timesteps
|
||||
sigmas = batch.sigmas
|
||||
n_tokens = batch.n_tokens
|
||||
|
||||
extra_set_timesteps_kwargs = {}
|
||||
if n_tokens is not None and "n_tokens" in inspect.signature(scheduler.set_timesteps).parameters:
|
||||
extra_set_timesteps_kwargs["n_tokens"] = n_tokens
|
||||
|
||||
# Flux2/BFL: Diffusers' Flux2 pipeline passes a custom sigma grid and
|
||||
# always supplies the resolution-dependent mu when the scheduler accepts
|
||||
# it.
|
||||
scheduler_config = getattr(scheduler, "config", None)
|
||||
use_flow_sigmas = (getattr(scheduler_config, "use_flow_sigmas", False) if scheduler_config else False)
|
||||
if timesteps is None and sigmas is None and not use_flow_sigmas:
|
||||
sigmas = np.linspace(1.0, 1.0 / num_inference_steps, num_inference_steps)
|
||||
|
||||
if "mu" in inspect.signature(scheduler.set_timesteps).parameters:
|
||||
if batch.n_tokens is not None:
|
||||
image_seq_len = batch.n_tokens
|
||||
else:
|
||||
h = (batch.height if isinstance(batch.height, int) else (batch.height[0] if batch.height else None))
|
||||
w = (batch.width if isinstance(batch.width, int) else (batch.width[0] if batch.width else None))
|
||||
vae_config = getattr(fastvideo_args.pipeline_config, "vae_config", None)
|
||||
if vae_config is not None:
|
||||
arch = getattr(vae_config, "arch_config", None)
|
||||
scale = (getattr(arch, "spatial_compression_ratio", 8) if arch else 8)
|
||||
else:
|
||||
scale = 8
|
||||
image_seq_len = ((h // scale) * (w // scale) if h is not None and w is not None else 256)
|
||||
extra_set_timesteps_kwargs["mu"] = compute_empirical_mu(image_seq_len, num_inference_steps)
|
||||
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. "
|
||||
"Please choose one to set custom values")
|
||||
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in inspect.signature(scheduler.set_timesteps).parameters
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(f"The current scheduler class {scheduler.__class__}'s "
|
||||
f"`set_timesteps` does not support custom timestep schedules.")
|
||||
timesteps_for_scheduler = (timesteps.cpu() if isinstance(timesteps, torch.Tensor) else timesteps)
|
||||
scheduler.set_timesteps(
|
||||
timesteps=timesteps_for_scheduler,
|
||||
device=device,
|
||||
**extra_set_timesteps_kwargs,
|
||||
)
|
||||
timesteps = scheduler.timesteps
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in inspect.signature(scheduler.set_timesteps).parameters
|
||||
if not accept_sigmas:
|
||||
raise ValueError(f"The current scheduler class {scheduler.__class__}'s "
|
||||
f"`set_timesteps` does not support custom sigmas schedules.")
|
||||
scheduler.set_timesteps(
|
||||
sigmas=sigmas,
|
||||
device=device,
|
||||
**extra_set_timesteps_kwargs,
|
||||
)
|
||||
timesteps = scheduler.timesteps
|
||||
else:
|
||||
scheduler.set_timesteps(
|
||||
num_inference_steps,
|
||||
device=device,
|
||||
**extra_set_timesteps_kwargs,
|
||||
)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
batch.timesteps = timesteps
|
||||
return batch
|
||||
@@ -0,0 +1,75 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Flux2 model family pipeline presets.
|
||||
|
||||
Each preset is a named inference preset that declares the user-facing
|
||||
stage topology, default sampling values, and which per-stage overrides
|
||||
are allowed. Presets are registered explicitly from
|
||||
:func:`fastvideo.registry._register_presets`.
|
||||
"""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Main denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
FLUX2_DEV = InferencePreset(
|
||||
name="flux2_dev",
|
||||
version=1,
|
||||
model_family="flux2",
|
||||
description="Flux2 full T2I with embedded guidance",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"seed": 0,
|
||||
"guidance_scale": 4.0,
|
||||
"num_inference_steps": 50,
|
||||
},
|
||||
)
|
||||
|
||||
FLUX2_KLEIN_4B = InferencePreset(
|
||||
name="flux2_klein_4b",
|
||||
version=1,
|
||||
model_family="flux2",
|
||||
description="Flux2 Klein 4B (distilled, 4-step, no guidance)",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"seed": 0,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
},
|
||||
)
|
||||
|
||||
FLUX2_KLEIN_9B = InferencePreset(
|
||||
name="flux2_klein_9b",
|
||||
version=1,
|
||||
model_family="flux2",
|
||||
description="Flux2 Klein 9B (distilled, 4-step, no guidance)",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"seed": 0,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (FLUX2_DEV, FLUX2_KLEIN_4B, FLUX2_KLEIN_9B)
|
||||
@@ -0,0 +1,80 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Lucy Edit video editing pipeline.
|
||||
|
||||
Lucy Edit uses a Wan2.2 5B transformer with an input video latent appended to
|
||||
the noisy latent channels. The stage topology is therefore closest to Wan V2V,
|
||||
but the model repo does not include CLIP image-encoder components.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.wan.wan_v2v_pipeline import WanVideoToVideoPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
VideoVAEEncodingStage,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LucyEditPipeline(WanVideoToVideoPipeline):
|
||||
"""FastVideo pipeline for decart-ai/Lucy-Edit-Dev."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
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"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_latent_preparation_stage",
|
||||
stage=VideoVAEEncodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = LucyEditPipeline
|
||||
@@ -268,6 +268,24 @@ FAST_WAN_2_2_TI2V_5B = InferencePreset(
|
||||
},
|
||||
)
|
||||
|
||||
LUCY_EDIT_DEV = InferencePreset(
|
||||
name="lucy_edit_dev",
|
||||
version=1,
|
||||
model_family="wan",
|
||||
description="Lucy Edit Dev 5B video editing",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 24,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": "",
|
||||
},
|
||||
)
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Self-Forcing (causal) presets
|
||||
# -------------------------------------------------------------------
|
||||
@@ -341,6 +359,7 @@ ALL_PRESETS = (
|
||||
FAST_WAN_T2V_480P,
|
||||
WAN_2_2_TI2V_5B,
|
||||
FAST_WAN_2_2_TI2V_5B,
|
||||
LUCY_EDIT_DEV,
|
||||
SF_WAN_T2V_1_3B,
|
||||
SF_WAN_2_2_T2V_A14B,
|
||||
SF_WAN_2_2_I2V_A14B,
|
||||
|
||||
@@ -380,6 +380,8 @@ class ComposedPipelineBase(ABC):
|
||||
model_index.pop("boundary_ratio", None)
|
||||
# used by Wan2.2 ti2v
|
||||
model_index.pop("expand_timesteps", None)
|
||||
# HF metadata (e.g. Flux2 Klein is_distilled); not a loadable module
|
||||
model_index.pop("is_distilled", None)
|
||||
|
||||
# some sanity checks
|
||||
assert len(model_index) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
@@ -47,6 +47,23 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def _is_flux2_packed(self, latents: torch.Tensor) -> bool:
|
||||
"""Detect Flux2 packed latents by checking channel count against VAE geometry.
|
||||
|
||||
Flux2 packs latent_channels into 2x2 spatial patches, so the DiT
|
||||
operates on ``latent_channels * 4`` channels at half spatial resolution.
|
||||
The VAE's ``post_quant_conv`` input dimension equals ``latent_channels``.
|
||||
"""
|
||||
if not hasattr(self.vae, "bn"):
|
||||
return False
|
||||
pqc = getattr(self.vae, "post_quant_conv", None)
|
||||
if pqc is None:
|
||||
return False
|
||||
vae_latent_ch = pqc.weight.shape[1]
|
||||
packed_ch = vae_latent_ch * 4 # 2x2 patch packing
|
||||
ch_dim = 1 if latents.ndim >= 4 else -1
|
||||
return latents.shape[ch_dim] == packed_ch
|
||||
|
||||
def _denormalize_latents(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""Convert normalized latents into the VAE's expected latent space."""
|
||||
# Some VAEs handle latent (de)normalization internally.
|
||||
@@ -77,6 +94,29 @@ class DecodingStage(PipelineStage):
|
||||
|
||||
return latents
|
||||
|
||||
@staticmethod
|
||||
def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor:
|
||||
"""Inverse of 2x2 patch packing: ``(B, C*4, H', W') -> (B, C, 2*H', 2*W')``."""
|
||||
batch_size, num_channels, height, width = latents.shape
|
||||
latents = latents.reshape(batch_size, num_channels // (2 * 2), 2, 2, height, width)
|
||||
latents = latents.permute(0, 1, 4, 2, 5, 3)
|
||||
latents = latents.reshape(batch_size, num_channels // (2 * 2), height * 2, width * 2)
|
||||
return latents
|
||||
|
||||
def _flux2_bn_denorm_and_unpatchify(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""BN denormalize then unpatchify packed latents for VAE decode.
|
||||
|
||||
Handles any channel count (e.g. 64->16, 128->32) via 2x2 spatial unpack.
|
||||
"""
|
||||
running_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype)
|
||||
running_var = self.vae.bn.running_var.view(1, -1, 1, 1).to(latents.device, latents.dtype)
|
||||
cfg = getattr(self.vae, "config", None)
|
||||
arch = getattr(cfg, "arch_config", None) if cfg else None
|
||||
eps = getattr(arch, "batch_norm_eps", None) or getattr(cfg, "batch_norm_eps", 1e-5)
|
||||
bn_std = torch.sqrt(torch.clamp(running_var + eps, min=1e-6))
|
||||
latents = latents * bn_std + running_mean
|
||||
return self._unpatchify_latents(latents)
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor, fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""
|
||||
@@ -100,7 +140,9 @@ class DecodingStage(PipelineStage):
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = self._denormalize_latents(latents)
|
||||
# Flux2: skip denormalize on packed latents; BN denorm runs below instead
|
||||
if not (latents.ndim == 5 and self._is_flux2_packed(latents)):
|
||||
latents = self._denormalize_latents(latents)
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
@@ -110,7 +152,27 @@ class DecodingStage(PipelineStage):
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
# Flux2's image VAE expects 4D (B, C, H, W); squeeze the singleton T
|
||||
# only for Flux2 packed latents. Gated on `_is_flux2_packed` so video
|
||||
# VAEs that legitimately decode 5D latents with T=1 are untouched.
|
||||
squeezed_for_vae = False
|
||||
if latents.ndim == 5 and latents.shape[2] == 1 and self._is_flux2_packed(latents):
|
||||
latents = latents.squeeze(2)
|
||||
squeezed_for_vae = True
|
||||
# Flux2 packed: BN denorm + unpatchify for VAE decode.
|
||||
# BN denorm is the complete inverse normalisation for Flux2 (no
|
||||
# scaling_factor/shift_factor step), matching Diffusers.
|
||||
if latents.ndim == 4 and self._is_flux2_packed(latents):
|
||||
latents = self._flux2_bn_denorm_and_unpatchify(latents)
|
||||
image = self.vae.decode(latents)
|
||||
# Unwrap diffusers-style DecoderOutput / tuple (Flux2 VAE returns a
|
||||
# DecoderOutput). No-op for existing VAEs that return a plain tensor.
|
||||
if hasattr(image, "sample"):
|
||||
image = image.sample
|
||||
elif isinstance(image, tuple | list):
|
||||
image = image[0]
|
||||
if squeezed_for_vae:
|
||||
image = image.unsqueeze(2)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
@@ -201,7 +263,13 @@ class DecodingStage(PipelineStage):
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
frames = batch.latents if fastvideo_args.output_type == "latent" else self.decode(batch.latents, fastvideo_args)
|
||||
if fastvideo_args.output_type == "latent":
|
||||
frames = batch.latents
|
||||
if frames.ndim == 5 and frames.shape[2] == 1 and self._is_flux2_packed(frames):
|
||||
frames = self._flux2_bn_denorm_and_unpatchify(frames.squeeze(2))
|
||||
frames = frames.unsqueeze(2)
|
||||
else:
|
||||
frames = self.decode(batch.latents, fastvideo_args)
|
||||
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
|
||||
@@ -4,6 +4,7 @@ Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import os
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
@@ -103,7 +104,37 @@ class DenoisingStage(PipelineStage):
|
||||
# TODO(will): make the precision configurable for inference
|
||||
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
# Flux2-only denoising compensations.
|
||||
#
|
||||
# `_is_flux` gates four behaviors that exist because Flux2's transformer
|
||||
# forward() does things internally that the generic pipeline must undo or
|
||||
# match. These are architectural facts about the Flux2 transformer, not
|
||||
# tunable precision policies (the precision policies #5/#6 — prompt-embed
|
||||
# casting and scheduler-step placement — were already moved to config:
|
||||
# DiTArchConfig.cast_prompt_embeds_to_dit_dtype and
|
||||
# PipelineConfig.scheduler_step_in_fp32).
|
||||
#
|
||||
# The four behaviors gated below:
|
||||
# 1. env-var bf16-reduced-precision matmul disable (4-step Klein drift)
|
||||
# 2. autocast disabled (Flux2 long-sequence attention breaks parity under autocast)
|
||||
# 3. guidance: skip the external x1000 (Flux2 multiplies guidance by 1000 internally)
|
||||
# 4. timestep: divide by 1000 with cast-before-divide (Flux2 multiplies timestep by 1000 internally)
|
||||
#
|
||||
# Contract: `prefix == "Flux"` is set ONLY by Flux2 (fastvideo/configs/
|
||||
# models/dits/flux_2.py). No other model uses that prefix, so this exact
|
||||
# match cannot false-positive. A future Flux variant that needs the same
|
||||
# compensations must either set prefix == "Flux" too, OR (preferred) these
|
||||
# gates should graduate to arch-config declarations like the precision
|
||||
# policies above.
|
||||
_is_flux = (getattr(fastvideo_args.pipeline_config.dit_config, "prefix", "") == "Flux")
|
||||
if _is_flux and os.getenv("FASTVIDEO_FLUX2_DISABLE_BF16_REDUCED_PRECISION_REDUCTION",
|
||||
"").lower() in {"1", "true", "yes"}:
|
||||
# Gate 1: tighten bf16 matmul accumulation for the 4-step Klein model (opt-in via env var).
|
||||
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False
|
||||
# Gate 2: Flux2 runs its bf16 transformer WITHOUT autocast — autocast perturbs long-sequence attention enough to break 4-step latent parity.
|
||||
autocast_enabled = ((target_dtype != torch.float32) and not fastvideo_args.disable_autocast and not _is_flux)
|
||||
scheduler_fp32 = getattr(fastvideo_args.pipeline_config, "scheduler_step_in_fp32", False)
|
||||
local_device = get_local_torch_device()
|
||||
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
@@ -159,13 +190,40 @@ class DenoisingStage(PipelineStage):
|
||||
},
|
||||
)
|
||||
|
||||
for key in ("flux2_txt_ids", "flux2_img_ids"):
|
||||
value = batch.extra.get(key)
|
||||
if torch.is_tensor(value):
|
||||
batch.extra[key] = value.to(device=local_device)
|
||||
|
||||
flux2_id_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"txt_ids": batch.extra.get("flux2_txt_ids"),
|
||||
"img_ids": batch.extra.get("flux2_img_ids"),
|
||||
},
|
||||
)
|
||||
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
cast_embeds = getattr(fastvideo_args.pipeline_config.dit_config, "cast_prompt_embeds_to_dit_dtype", False)
|
||||
if cast_embeds:
|
||||
prompt_embeds = [
|
||||
embed.to(device=local_device, dtype=target_dtype) if torch.is_tensor(embed) else embed
|
||||
for embed in batch.prompt_embeds
|
||||
]
|
||||
else:
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert not torch.isnan(prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert neg_prompt_embeds is not None
|
||||
if cast_embeds:
|
||||
neg_prompt_embeds = [
|
||||
embed.to(device=local_device, dtype=target_dtype) if torch.is_tensor(embed) else embed
|
||||
for embed in neg_prompt_embeds
|
||||
]
|
||||
else:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert not torch.isnan(neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
@@ -212,23 +270,34 @@ class DenoisingStage(PipelineStage):
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
is_lucy_edit = fastvideo_args.pipeline_config.lucy_edit_task
|
||||
|
||||
# Hoisted out of the per-step loop: depends only on inputs that
|
||||
# are constant across denoising steps.
|
||||
use_meanflow = getattr(self.transformer.config, "use_meanflow", False)
|
||||
# Gate 3: Flux2's transformer multiplies guidance by 1000 internally, so we
|
||||
# skip the external *1000 pre-scaling for Flux models.
|
||||
embedded_cfg_scale = fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
if _is_flux and embedded_cfg_scale is not None:
|
||||
embedded_cfg_scale = batch.guidance_scale
|
||||
if embedded_cfg_scale is not None:
|
||||
guidance_expand = (torch.tensor(
|
||||
[embedded_cfg_scale] * latents.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) * 1000.0)
|
||||
).to(target_dtype) * (1.0 if _is_flux else 1000.0))
|
||||
else:
|
||||
guidance_expand = None
|
||||
# V2V padding: zero-filled tensor concatenated with each step's
|
||||
# latent_model_input. Shape is fixed by latents and is never
|
||||
# written to, so we allocate once.
|
||||
v2v_zero_pad = torch.zeros_like(latents) if batch.video_latent is not None else None
|
||||
lucy_timestep_seq_len = None
|
||||
if is_lucy_edit:
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
assert patch_size[0] == 1, "Lucy Edit timestep expansion assumes temporal patch size 1"
|
||||
lucy_timestep_seq_len = (latents.shape[2] * (latents.shape[3] // patch_size[1]) *
|
||||
(latents.shape[4] // patch_size[2]))
|
||||
|
||||
# CFG gating / stale-uncond reuse setup (Adaptive Guidance LinearAG
|
||||
# variant, Castillo et al. 2023). When envs.FASTVIDEO_CFG_GATE_STEP
|
||||
@@ -309,14 +378,25 @@ class DenoisingStage(PipelineStage):
|
||||
# Expand latents for V2V/I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.video_latent is not None:
|
||||
latent_model_input = torch.cat([latent_model_input, batch.video_latent, v2v_zero_pad],
|
||||
dim=1).to(target_dtype)
|
||||
if is_lucy_edit:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.video_latent],
|
||||
dim=1,
|
||||
).to(target_dtype)
|
||||
else:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.video_latent, v2v_zero_pad],
|
||||
dim=1,
|
||||
).to(target_dtype)
|
||||
elif batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat([latent_model_input, batch.image_latent], dim=1).to(target_dtype)
|
||||
|
||||
assert not torch.isnan(latent_model_input).any(), "latent_model_input contains nan"
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
if is_lucy_edit:
|
||||
assert lucy_timestep_seq_len is not None
|
||||
t_expand = t.repeat(latent_model_input.shape[0], lucy_timestep_seq_len)
|
||||
elif fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
|
||||
temp_ts = torch.cat([temp_ts, temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep])
|
||||
@@ -324,7 +404,19 @@ class DenoisingStage(PipelineStage):
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
t_expand = t_expand.to(get_local_torch_device())
|
||||
# Gate 4: Flux2 transformer multiplies timestep by 1000 internally, so
|
||||
# the pipeline must pass timestep/1000 (matching Diffusers).
|
||||
# Diffusers casts to the latent dtype before the division; doing
|
||||
# the division in fp32 first changes BF16 rounding for the final
|
||||
# Klein timestep and breaks latent parity.
|
||||
if _is_flux:
|
||||
t_expand = t_expand.to(
|
||||
device=get_local_torch_device(),
|
||||
dtype=latent_model_input.dtype,
|
||||
)
|
||||
t_expand = t_expand / 1000.0
|
||||
else:
|
||||
t_expand = t_expand.to(get_local_torch_device())
|
||||
|
||||
if use_meanflow:
|
||||
if i == len(timesteps) - 1:
|
||||
@@ -403,6 +495,7 @@ class DenoisingStage(PipelineStage):
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
@@ -444,6 +537,7 @@ class DenoisingStage(PipelineStage):
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
_cfg_gate_fresh_uncond += 1
|
||||
|
||||
@@ -467,12 +561,16 @@ class DenoisingStage(PipelineStage):
|
||||
noise_pred_text,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
# Compute the previous noisy sample
|
||||
if scheduler_fp32:
|
||||
# Diffusers-style: fp32 Euler update outside autocast avoids BF16 drift.
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
else:
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# save trajectory latents if needed
|
||||
if batch.return_trajectory_latents:
|
||||
|
||||
@@ -629,9 +629,10 @@ class VideoVAEEncodingStage(ImageVAEEncodingStage):
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
sample_mode = "argmax" if fastvideo_args.pipeline_config.lucy_edit_task else "sample"
|
||||
if sample_mode == "sample" and generator is None:
|
||||
raise ValueError("Generator must be provided for sampled video VAE encoding")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator, sample_mode=sample_mode)
|
||||
|
||||
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
|
||||
@@ -57,6 +57,10 @@ class TextEncodingStage(PipelineStage):
|
||||
assert len(self.tokenizers) == len(self.text_encoders)
|
||||
assert len(self.text_encoders) == len(fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
|
||||
# Skip encoding if precomputed prompt_embeds were provided
|
||||
if batch.prompt_embeds is not None and len(batch.prompt_embeds) > 0:
|
||||
return batch
|
||||
|
||||
# Encode positive prompt with all available encoders
|
||||
assert batch.prompt is not None
|
||||
prompt_text: str | list[str] = batch.prompt
|
||||
@@ -218,7 +222,8 @@ class TextEncodingStage(PipelineStage):
|
||||
# Qwen2-style tokenizers. Scoped via treat_empty_as_dot so
|
||||
# models that legitimately use "" (e.g. negative_prompt="")
|
||||
# are not affected.
|
||||
if not processed_text.strip() and getattr(encoder_config, "treat_empty_as_dot", False):
|
||||
if isinstance(processed_text, str) and not processed_text.strip() and getattr(
|
||||
encoder_config, "treat_empty_as_dot", False):
|
||||
processed_text = "."
|
||||
processed_texts.append(processed_text)
|
||||
else:
|
||||
@@ -235,7 +240,27 @@ class TextEncodingStage(PipelineStage):
|
||||
tok = getattr(tokenizer, "tokenizer", tokenizer)
|
||||
|
||||
if encoder_config.is_chat_model:
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(target_device)
|
||||
already_chat_formatted = bool(processed_texts) and isinstance(processed_texts[0], list)
|
||||
if already_chat_formatted:
|
||||
# Existing chat models (e.g. HunyuanVideo 1.5 / Qwen2.5-VL)
|
||||
# pre-format prompts into message lists upstream and rely on
|
||||
# the inner tokenizer + full tokenizer_kwargs (which include
|
||||
# add_generation_prompt). Preserve that original path exactly.
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(target_device)
|
||||
else:
|
||||
# Two-step approach matching Diffusers: format with chat
|
||||
# template first, then tokenize the resulting strings.
|
||||
formatted_texts = []
|
||||
for pt in processed_texts:
|
||||
messages = [{"role": "user", "content": pt}]
|
||||
formatted = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
formatted_texts.append(formatted)
|
||||
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(target_device)
|
||||
else:
|
||||
text_inputs = tok(processed_texts, **tok_kwargs).to(target_device)
|
||||
|
||||
|
||||
+69
-9
@@ -30,6 +30,10 @@ from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.flux_2 import (
|
||||
Flux2KleinPipelineConfig,
|
||||
Flux2PipelineConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.configs.pipelines.turbodiffusion import (
|
||||
@@ -40,6 +44,7 @@ from fastvideo.configs.pipelines.turbodiffusion import (
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config,
|
||||
FastWan2_2_TI2V_5B_Config,
|
||||
LucyEditDevConfig,
|
||||
SelfForcingWan2_2_T2V480PConfig,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
WANV2VConfig,
|
||||
@@ -324,6 +329,45 @@ def _register_configs() -> None:
|
||||
default_preset="stable_audio_open_small",
|
||||
)
|
||||
|
||||
def _is_flux2_klein(path: str) -> bool:
|
||||
path_lower = path.lower()
|
||||
return "flux.2-klein" in path_lower or "flux2-klein" in path_lower or "flux2klein" in path_lower
|
||||
|
||||
def _is_flux2_full(path: str) -> bool:
|
||||
path_lower = path.lower()
|
||||
is_flux2 = "flux.2" in path_lower or "flux2" in path_lower or "flux_2" in path_lower or "flux-2" in path_lower
|
||||
return is_flux2 and "klein" not in path_lower
|
||||
|
||||
# Flux2 Klein (distilled, 4-step, no guidance)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Flux2KleinPipelineConfig,
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
hf_model_paths=[
|
||||
"black-forest-labs/FLUX.2-klein-4B",
|
||||
"black-forest-labs/FLUX.2-klein-9B",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_flux2_klein,
|
||||
],
|
||||
model_family="flux2",
|
||||
default_preset="flux2_klein_4b",
|
||||
)
|
||||
# Flux2 (full, Mistral3 text encoder, embedded guidance)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Flux2PipelineConfig,
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
hf_model_paths=[
|
||||
"black-forest-labs/FLUX.2-dev",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_flux2_full,
|
||||
],
|
||||
model_family="flux2",
|
||||
default_preset="flux2_dev",
|
||||
)
|
||||
|
||||
# Hunyuan 1.5 (specific)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -738,6 +782,18 @@ def _register_configs() -> None:
|
||||
model_family="wan",
|
||||
default_preset="fast_wan_2_2_ti2v_5b",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LucyEditDevConfig,
|
||||
workload_types=(),
|
||||
hf_model_paths=[
|
||||
"decart-ai/Lucy-Edit-Dev",
|
||||
"decart-ai/Lucy-Edit-1.1-Dev",
|
||||
],
|
||||
model_detectors=[lambda path: "lucy-edit" in path.lower()],
|
||||
model_family="wan",
|
||||
default_preset="lucy_edit_dev",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Wan2_2_T2V_A14B_Config,
|
||||
@@ -840,15 +896,19 @@ def get_model_info(
|
||||
if workload_type is None:
|
||||
workload_type = WorkloadType.T2V
|
||||
|
||||
if os.path.exists(model_path):
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
else:
|
||||
config = maybe_download_model_index(model_path)
|
||||
config_info = _get_config_info(model_path, raise_on_missing=True)
|
||||
assert config_info is not None, "config_info must be resolved"
|
||||
|
||||
pipeline_name = config.get("_class_name")
|
||||
if override_pipeline_cls_name:
|
||||
logger.info("Overriding pipeline class name from %s to %s", pipeline_name, override_pipeline_cls_name)
|
||||
pipeline_name = override_pipeline_cls_name
|
||||
logger.info("Using override pipeline class name %s", pipeline_name)
|
||||
else:
|
||||
if os.path.exists(model_path):
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
else:
|
||||
config = maybe_download_model_index(model_path)
|
||||
|
||||
pipeline_name = config.get("_class_name")
|
||||
|
||||
if pipeline_name is None:
|
||||
raise ValueError("Model config does not contain a _class_name attribute. "
|
||||
@@ -857,9 +917,6 @@ def get_model_info(
|
||||
pipeline_registry = get_pipeline_registry(pipeline_type)
|
||||
pipeline_cls = pipeline_registry.resolve_pipeline_cls(pipeline_name, pipeline_type, workload_type)
|
||||
|
||||
config_info = _get_config_info(model_path, raise_on_missing=True)
|
||||
assert config_info is not None, "config_info must be resolved"
|
||||
|
||||
sampling_param_cls = config_info.sampling_param_cls or SamplingParam
|
||||
|
||||
return ModelInfo(
|
||||
@@ -920,9 +977,12 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as TURBODIFFUSION_PRESETS, )
|
||||
from fastvideo.pipelines.basic.wan.presets import (
|
||||
ALL_PRESETS as WAN_PRESETS, )
|
||||
from fastvideo.pipelines.basic.flux_2.presets import (
|
||||
ALL_PRESETS as FLUX2_PRESETS, )
|
||||
|
||||
all_preset_groups = (
|
||||
COSMOS_PRESETS,
|
||||
FLUX2_PRESETS,
|
||||
GAMECRAFT_PRESETS,
|
||||
GEN3C_PRESETS,
|
||||
HUNYUAN_PRESETS,
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.configs.pipelines.wan import LucyEditDevConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.pipelines.basic.wan.lucy_edit_pipeline import LucyEditPipeline
|
||||
from fastvideo.pipelines.pipeline_registry import PipelineType, get_pipeline_registry
|
||||
from fastvideo.registry import get_default_preset, get_pipeline_config_cls_from_name
|
||||
|
||||
|
||||
def test_lucy_edit_registry_and_preset() -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
|
||||
preset = get_preset("lucy_edit_dev", "wan")
|
||||
assert preset.model_family == "wan"
|
||||
assert preset.defaults["height"] == 480
|
||||
assert preset.defaults["width"] == 832
|
||||
assert preset.defaults["num_frames"] == 81
|
||||
|
||||
config = LucyEditDevConfig()
|
||||
assert config.lucy_edit_task is True
|
||||
assert config.ti2v_task is False
|
||||
assert config.dit_config.arch_config.out_channels == 48
|
||||
assert config.dit_config.arch_config.in_channels == 96
|
||||
assert config.vae_config.arch_config.z_dim == 48
|
||||
assert config.dit_config.arch_config.in_channels == config.vae_config.arch_config.z_dim * 2
|
||||
assert get_default_preset("decart-ai/Lucy-Edit-Dev") == "lucy_edit_dev"
|
||||
assert get_default_preset("decart-ai/Lucy-Edit-1.1-Dev") == "lucy_edit_dev"
|
||||
assert get_pipeline_config_cls_from_name("decart-ai/Lucy-Edit-Dev") is LucyEditDevConfig
|
||||
assert get_pipeline_config_cls_from_name("decart-ai/Lucy-Edit-1.1-Dev") is LucyEditDevConfig
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("decart-ai/Lucy-Edit-Dev")
|
||||
assert sampling_param.height == 480
|
||||
assert sampling_param.width == 832
|
||||
assert sampling_param.num_frames == 81
|
||||
assert sampling_param.fps == 24
|
||||
assert sampling_param.guidance_scale == 5.0
|
||||
assert sampling_param.negative_prompt == ""
|
||||
|
||||
sampling_param_1_1 = SamplingParam.from_pretrained("decart-ai/Lucy-Edit-1.1-Dev")
|
||||
assert sampling_param_1_1.height == 480
|
||||
assert sampling_param_1_1.width == 832
|
||||
assert sampling_param_1_1.num_frames == 81
|
||||
assert sampling_param_1_1.fps == 24
|
||||
assert sampling_param_1_1.guidance_scale == 5.0
|
||||
assert sampling_param_1_1.negative_prompt == ""
|
||||
|
||||
# FastVideo has no V2V workload enum today; model_index dispatches Lucy by pipeline class name.
|
||||
registry = get_pipeline_registry(PipelineType.BASIC)
|
||||
assert registry.resolve_pipeline_cls("LucyEditPipeline", PipelineType.BASIC, WorkloadType.T2V) is LucyEditPipeline
|
||||
@@ -0,0 +1,38 @@
|
||||
import torch
|
||||
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionImpl,
|
||||
VideoSparseAttentionMetadataBuilder,
|
||||
)
|
||||
|
||||
|
||||
def _build_metadata(cache_tile_buf: bool):
|
||||
return VideoSparseAttentionMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=(4, 4, 4),
|
||||
patch_size=(1, 1, 1),
|
||||
VSA_sparsity=0.5,
|
||||
device=torch.device("cpu"),
|
||||
cache_tile_buf=cache_tile_buf,
|
||||
)
|
||||
|
||||
|
||||
def test_vsa_tile_does_not_cache_training_scratch_when_disabled():
|
||||
metadata = _build_metadata(cache_tile_buf=False)
|
||||
impl = object.__new__(VideoSparseAttentionImpl)
|
||||
x = torch.ones(1, 64, 2, 2)
|
||||
|
||||
tiled = impl.tile(x, metadata)
|
||||
|
||||
assert tiled.shape == x.shape
|
||||
assert metadata.tile_buf is None
|
||||
|
||||
|
||||
def test_vsa_tile_caches_scratch_by_default():
|
||||
metadata = _build_metadata(cache_tile_buf=True)
|
||||
impl = object.__new__(VideoSparseAttentionImpl)
|
||||
x = torch.ones(1, 64, 2, 2)
|
||||
|
||||
tiled = impl.tile(x, metadata)
|
||||
|
||||
assert metadata.tile_buf is tiled
|
||||
@@ -0,0 +1,49 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Regression coverage for PR #1390's S2-1 plumbing finding: the dormant FP4 shape-tracking path must stay
|
||||
gated off unless explicitly enabled, and the MLP quant_config=None default path must keep using ReplicatedLinear's
|
||||
unquantized fallback.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.layers.linear import ReplicatedLinear, UnquantizedLinearMethod
|
||||
from fastvideo.layers.mlp import MLP
|
||||
|
||||
|
||||
def test_replicated_linear_shape_tracking_default_off() -> None:
|
||||
ReplicatedLinear.reset_shape_tracking()
|
||||
assert ReplicatedLinear.enable_shape_tracking is False
|
||||
|
||||
linear = ReplicatedLinear(input_size=8, output_size=4)
|
||||
linear(torch.randn(2, 8))
|
||||
|
||||
assert len(ReplicatedLinear._shape_to_layer_types) == 0
|
||||
|
||||
|
||||
def test_replicated_linear_shape_tracking_enabled_records_unique_shapes() -> None:
|
||||
ReplicatedLinear.reset_shape_tracking()
|
||||
ReplicatedLinear.enable_shape_tracking = True
|
||||
try:
|
||||
linear = ReplicatedLinear(input_size=8, output_size=4)
|
||||
linear(torch.randn(2, 8))
|
||||
linear(torch.randn(3, 8))
|
||||
|
||||
assert len(ReplicatedLinear._shape_to_layer_types) == 2
|
||||
for layer_types in ReplicatedLinear._shape_to_layer_types.values():
|
||||
assert "ReplicatedLinear" in layer_types
|
||||
|
||||
ReplicatedLinear.reset_shape_tracking()
|
||||
assert len(ReplicatedLinear._shape_to_layer_types) == 0
|
||||
finally:
|
||||
ReplicatedLinear.enable_shape_tracking = False
|
||||
|
||||
|
||||
def test_mlp_quant_config_none_uses_unquantized_path() -> None:
|
||||
mlp = MLP(input_dim=8, mlp_hidden_dim=16)
|
||||
|
||||
assert isinstance(mlp.fc_in.quant_method, UnquantizedLinearMethod)
|
||||
assert isinstance(mlp.fc_out.quant_method, UnquantizedLinearMethod)
|
||||
|
||||
output = mlp.forward(torch.randn(2, 8))
|
||||
assert output.shape == (2, 8)
|
||||
@@ -0,0 +1,531 @@
|
||||
"""Launch an arbitrary FastVideo command on Modal GPUs.
|
||||
|
||||
Examples:
|
||||
python -m modal run fastvideo/tests/modal/launch_l40s_job.py --command "nvidia-smi" --install-extra none
|
||||
|
||||
python -m modal run fastvideo/tests/modal/launch_l40s_job.py \
|
||||
--num-gpus 2 \
|
||||
--install-extra test \
|
||||
--command "pytest fastvideo/tests/vaes -vs"
|
||||
|
||||
python -m modal run fastvideo/tests/modal/launch_l40s_job.py \
|
||||
--gpu-type H100 \
|
||||
--num-gpus 1 \
|
||||
--install-extra none \
|
||||
--command "nvidia-smi"
|
||||
|
||||
Use ``--no-wait`` with ``modal run --detach`` when the job should keep running
|
||||
after the local Modal client exits.
|
||||
"""
|
||||
|
||||
import os
|
||||
import base64
|
||||
import shlex
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import modal
|
||||
|
||||
app = modal.App("fastvideo-gpu-job")
|
||||
|
||||
REPO_DIR = "/FastVideo"
|
||||
MODEL_VOLUME_NAME = os.environ.get("FASTVIDEO_MODAL_VOLUME", "hf-model-weights")
|
||||
IMAGE_VERSION = os.environ.get("IMAGE_VERSION", "latest")
|
||||
IMAGE_TAG = os.environ.get(
|
||||
"FASTVIDEO_MODAL_IMAGE",
|
||||
f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{IMAGE_VERSION}",
|
||||
)
|
||||
SECRET_ENV_KEYS = (
|
||||
"HF_API_KEY",
|
||||
"HUGGINGFACE_HUB_TOKEN",
|
||||
"HF_TOKEN",
|
||||
"WANDB_API_KEY",
|
||||
"WANDB_BASE_URL",
|
||||
"WANDB_MODE",
|
||||
)
|
||||
|
||||
print(f"Using image: {IMAGE_TAG}")
|
||||
print(f"Using Modal volume: {MODEL_VOLUME_NAME}")
|
||||
|
||||
model_vol = modal.Volume.from_name(MODEL_VOLUME_NAME, create_if_missing=True)
|
||||
local_secrets = modal.Secret.from_dict({
|
||||
key: os.environ[key]
|
||||
for key in SECRET_ENV_KEYS
|
||||
if os.environ.get(key)
|
||||
})
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(IMAGE_TAG, add_python="3.12")
|
||||
.apt_install(
|
||||
"cmake",
|
||||
"pkg-config",
|
||||
"build-essential",
|
||||
"curl",
|
||||
"git",
|
||||
"libssl-dev",
|
||||
"ffmpeg",
|
||||
)
|
||||
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
|
||||
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
|
||||
.env({
|
||||
"PATH": "/root/.cargo/bin:$PATH",
|
||||
"HF_HOME": "/root/data/.cache",
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
})
|
||||
)
|
||||
|
||||
COMMON_FUNCTION_KWARGS = dict(
|
||||
image=image,
|
||||
timeout=86400,
|
||||
secrets=[local_secrets],
|
||||
volumes={"/root/data": model_vol},
|
||||
)
|
||||
|
||||
|
||||
def _run_local_git_command(args: list[str]) -> str:
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def _run_local_git_command_allow_diff(args: list[str]) -> str:
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode not in (0, 1):
|
||||
raise RuntimeError(result.stderr.strip() or f"git {' '.join(args)} failed")
|
||||
return result.stdout
|
||||
|
||||
|
||||
def _split_patch_paths(patch_paths: str) -> list[str]:
|
||||
return [path.strip() for path in patch_paths.split(",") if path.strip()]
|
||||
|
||||
|
||||
def _build_local_patch(patch_paths: str) -> str:
|
||||
paths = _split_patch_paths(patch_paths)
|
||||
diff_args = ["diff", "--binary"]
|
||||
if paths:
|
||||
diff_args.extend(["--", *paths])
|
||||
patch_parts = [_run_local_git_command_allow_diff(diff_args)]
|
||||
|
||||
untracked_args = ["ls-files", "--others", "--exclude-standard"]
|
||||
if paths:
|
||||
untracked_args.extend(["--", *paths])
|
||||
untracked = _run_local_git_command(untracked_args).splitlines()
|
||||
for path in untracked:
|
||||
if not os.path.isfile(path):
|
||||
continue
|
||||
patch_parts.append(_run_local_git_command_allow_diff(["diff", "--binary", "--no-index", "/dev/null", path]))
|
||||
|
||||
patch = "\n".join(part for part in patch_parts if part.strip())
|
||||
if not patch.strip():
|
||||
raise RuntimeError("Requested --apply-local-patch but no local diff was found.")
|
||||
return patch
|
||||
|
||||
|
||||
def _apply_local_patch(patch_b64: str) -> None:
|
||||
if not patch_b64:
|
||||
return
|
||||
patch = base64.b64decode(patch_b64.encode("ascii"))
|
||||
print("Applying local workspace patch", flush=True)
|
||||
result = subprocess.run(
|
||||
["git", "apply", "--binary", "--whitespace=nowarn", "-"],
|
||||
cwd=REPO_DIR,
|
||||
input=patch,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=sys.stderr,
|
||||
check=False,
|
||||
)
|
||||
if result.stdout:
|
||||
print(result.stdout.decode("utf-8", errors="replace"), end="", flush=True)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Failed to apply local patch with exit code {result.returncode}")
|
||||
|
||||
|
||||
def _normalize_git_repo_url(git_repo: str) -> str:
|
||||
if git_repo.startswith("git@github.com:"):
|
||||
return "https://github.com/" + git_repo[len("git@github.com:"):]
|
||||
if git_repo.startswith("ssh://git@github.com/"):
|
||||
return "https://github.com/" + git_repo[len("ssh://git@github.com/"):]
|
||||
return git_repo
|
||||
|
||||
|
||||
def _resolve_git_repo(git_repo: str) -> str:
|
||||
if git_repo.strip():
|
||||
return _normalize_git_repo_url(git_repo.strip())
|
||||
|
||||
env_repo = os.environ.get("BUILDKITE_REPO", "").strip()
|
||||
if env_repo:
|
||||
return _normalize_git_repo_url(env_repo)
|
||||
|
||||
discovered_repo = _run_local_git_command(["config", "--get", "remote.origin.url"])
|
||||
if discovered_repo:
|
||||
return _normalize_git_repo_url(discovered_repo)
|
||||
|
||||
raise RuntimeError("Could not resolve git repo URL. Pass --git-repo or set BUILDKITE_REPO.")
|
||||
|
||||
|
||||
def _resolve_git_commit(git_commit: str) -> str:
|
||||
if git_commit.strip():
|
||||
return git_commit.strip()
|
||||
|
||||
env_commit = os.environ.get("BUILDKITE_COMMIT", "").strip()
|
||||
if env_commit:
|
||||
return env_commit
|
||||
|
||||
discovered_commit = _run_local_git_command(["rev-parse", "HEAD"])
|
||||
if discovered_commit:
|
||||
return discovered_commit
|
||||
|
||||
raise RuntimeError("Could not resolve git commit. Pass --git-commit or set BUILDKITE_COMMIT.")
|
||||
|
||||
|
||||
def _resolve_pull_request(pr_number: str) -> str:
|
||||
if pr_number.strip():
|
||||
return pr_number.strip()
|
||||
env_pr = os.environ.get("BUILDKITE_PULL_REQUEST", "").strip()
|
||||
if env_pr:
|
||||
return env_pr
|
||||
return "false"
|
||||
|
||||
|
||||
def _run(args: list[str], cwd: str | None = None, env: dict[str, str] | None = None) -> str:
|
||||
print("$ " + " ".join(shlex.quote(arg) for arg in args), flush=True)
|
||||
result = subprocess.run(
|
||||
args,
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=sys.stderr,
|
||||
text=True,
|
||||
)
|
||||
if result.stdout:
|
||||
print(result.stdout, end="", flush=True)
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def _run_shell(command: str, cwd: str, env: dict[str, str]) -> None:
|
||||
print(f"$ {command}", flush=True)
|
||||
result = subprocess.run(
|
||||
["/bin/bash", "-lc", command],
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
stdout=sys.stdout,
|
||||
stderr=sys.stderr,
|
||||
check=False,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Command failed with exit code {result.returncode}: {command}")
|
||||
|
||||
|
||||
def _parse_env_vars(env_vars: str) -> dict[str, str]:
|
||||
parsed: dict[str, str] = {}
|
||||
for item in env_vars.split(","):
|
||||
item = item.strip()
|
||||
if not item:
|
||||
continue
|
||||
if "=" not in item:
|
||||
raise RuntimeError(f"Invalid env var override {item!r}; expected KEY=VALUE.")
|
||||
key, value = item.split("=", 1)
|
||||
parsed[key.strip()] = value.strip()
|
||||
return parsed
|
||||
|
||||
|
||||
def _activate_remote_python_env(env: dict[str, str]) -> dict[str, str]:
|
||||
venv_bin = "/opt/venv/bin"
|
||||
if os.path.isdir(venv_bin):
|
||||
env["VIRTUAL_ENV"] = "/opt/venv"
|
||||
env["PATH"] = venv_bin + os.pathsep + env.get("PATH", "")
|
||||
return env
|
||||
|
||||
|
||||
def _clone_checkout(git_repo: str, git_commit: str, pr_number: str) -> str:
|
||||
last_clone_error: subprocess.CalledProcessError | None = None
|
||||
for attempt in range(1, 4):
|
||||
shutil.rmtree(REPO_DIR, ignore_errors=True)
|
||||
try:
|
||||
_run(
|
||||
[
|
||||
"git",
|
||||
"-c",
|
||||
"http.version=HTTP/1.1",
|
||||
"clone",
|
||||
git_repo,
|
||||
REPO_DIR,
|
||||
],
|
||||
cwd="/",
|
||||
)
|
||||
break
|
||||
except subprocess.CalledProcessError as error:
|
||||
last_clone_error = error
|
||||
if attempt == 3:
|
||||
raise
|
||||
sleep_seconds = 5 * attempt
|
||||
print(
|
||||
f"git clone failed on attempt {attempt}; retrying in {sleep_seconds}s",
|
||||
flush=True,
|
||||
)
|
||||
time.sleep(sleep_seconds)
|
||||
if last_clone_error is not None and not os.path.isdir(REPO_DIR):
|
||||
raise last_clone_error
|
||||
if pr_number and pr_number != "false":
|
||||
_run(["git", "fetch", "--prune", "origin", f"refs/pull/{pr_number}/head"], cwd=REPO_DIR)
|
||||
_run(["git", "checkout", "FETCH_HEAD"], cwd=REPO_DIR)
|
||||
else:
|
||||
_run(["git", "checkout", git_commit], cwd=REPO_DIR)
|
||||
_run(["git", "submodule", "update", "--init", "--recursive"], cwd=REPO_DIR)
|
||||
return _run(["git", "rev-parse", "HEAD"], cwd=REPO_DIR)
|
||||
|
||||
|
||||
def _install_fastvideo(install_extra: str, env: dict[str, str]) -> None:
|
||||
install_extra = install_extra.strip()
|
||||
if install_extra.lower() in {"", "none", "skip", "false"}:
|
||||
return
|
||||
package = "." if install_extra == "." else f".[{install_extra}]"
|
||||
_run_shell(
|
||||
"source $HOME/.local/bin/env 2>/dev/null || true; "
|
||||
"source /opt/venv/bin/activate 2>/dev/null || true; "
|
||||
f"uv pip install -e {shlex.quote(package)}",
|
||||
cwd=REPO_DIR,
|
||||
env=env,
|
||||
)
|
||||
|
||||
|
||||
def _build_kernel(env: dict[str, str]) -> None:
|
||||
_run_shell(
|
||||
"source $HOME/.local/bin/env 2>/dev/null || true; "
|
||||
"source /opt/venv/bin/activate 2>/dev/null || true; "
|
||||
"./build.sh",
|
||||
cwd=os.path.join(REPO_DIR, "fastvideo-kernel"),
|
||||
env=env,
|
||||
)
|
||||
|
||||
|
||||
def _run_gpu_job(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
) -> dict[str, Any]:
|
||||
remote_env = _activate_remote_python_env(os.environ.copy())
|
||||
remote_env.update({
|
||||
"HF_HOME": "/root/data/.cache",
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
"FASTVIDEO_ATTENTION_BACKEND": remote_env.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
})
|
||||
remote_env.update(_parse_env_vars(env_vars))
|
||||
|
||||
print(f"Cloning repository: {git_repo}")
|
||||
print(f"Target commit: {git_commit}")
|
||||
if pr_number and pr_number != "false":
|
||||
print(f"Using PR ref: {pr_number}")
|
||||
checked_out_commit = _clone_checkout(git_repo, git_commit, pr_number)
|
||||
print(f"Checked out commit: {checked_out_commit}")
|
||||
_apply_local_patch(local_patch_b64)
|
||||
|
||||
_install_fastvideo(install_extra, remote_env)
|
||||
if build_kernel:
|
||||
_build_kernel(remote_env)
|
||||
|
||||
try:
|
||||
_run_shell(command, cwd=REPO_DIR, env=remote_env)
|
||||
finally:
|
||||
if commit_volume:
|
||||
print("Committing Modal volume", flush=True)
|
||||
model_vol.commit()
|
||||
return {
|
||||
"command": command,
|
||||
"git_repo": git_repo,
|
||||
"git_commit": checked_out_commit,
|
||||
"install_extra": install_extra,
|
||||
"build_kernel": build_kernel,
|
||||
"local_patch_applied": bool(local_patch_b64),
|
||||
"commit_volume": commit_volume,
|
||||
}
|
||||
|
||||
|
||||
@app.function(gpu="L40S:1", **COMMON_FUNCTION_KWARGS)
|
||||
def run_l40s_1(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
@app.function(gpu="L40S:2", **COMMON_FUNCTION_KWARGS)
|
||||
def run_l40s_2(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
@app.function(gpu="L40S:4", **COMMON_FUNCTION_KWARGS)
|
||||
def run_l40s_4(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
@app.function(gpu="L40S:8", **COMMON_FUNCTION_KWARGS)
|
||||
def run_l40s_8(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
@app.function(gpu="H100:1", **COMMON_FUNCTION_KWARGS)
|
||||
def run_h100_1(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
@app.function(gpu="H100:2", **COMMON_FUNCTION_KWARGS)
|
||||
def run_h100_2(
|
||||
command: str,
|
||||
git_repo: str,
|
||||
git_commit: str,
|
||||
pr_number: str,
|
||||
install_extra: str,
|
||||
build_kernel: bool,
|
||||
env_vars: str,
|
||||
local_patch_b64: str,
|
||||
commit_volume: bool,
|
||||
):
|
||||
return _run_gpu_job(command, git_repo, git_commit, pr_number, install_extra, build_kernel, env_vars,
|
||||
local_patch_b64, commit_volume)
|
||||
|
||||
|
||||
def _select_runner(gpu_type: str, num_gpus: int) -> Callable[..., Any]:
|
||||
normalized_gpu_type = gpu_type.upper()
|
||||
runners = {
|
||||
("L40S", 1): run_l40s_1,
|
||||
("L40S", 2): run_l40s_2,
|
||||
("L40S", 4): run_l40s_4,
|
||||
("L40S", 8): run_l40s_8,
|
||||
("H100", 1): run_h100_1,
|
||||
("H100", 2): run_h100_2,
|
||||
}
|
||||
try:
|
||||
return runners[(normalized_gpu_type, num_gpus)]
|
||||
except KeyError as error:
|
||||
supported = ", ".join(f"{gpu}:{count}" for gpu, count in sorted(runners))
|
||||
raise RuntimeError(f"Unsupported GPU request {gpu_type}:{num_gpus}. Supported requests: {supported}.") from error
|
||||
|
||||
|
||||
@app.local_entrypoint()
|
||||
def main(
|
||||
command: str = "nvidia-smi",
|
||||
gpu_type: str = "L40S",
|
||||
num_gpus: int = 1,
|
||||
git_repo: str = "",
|
||||
git_commit: str = "",
|
||||
pr_number: str = "",
|
||||
install_extra: str = "dev",
|
||||
build_kernel: bool = False,
|
||||
env_vars: str = "",
|
||||
apply_local_patch: bool = False,
|
||||
patch_paths: str = "",
|
||||
wait: bool = True,
|
||||
commit_volume: bool = False,
|
||||
):
|
||||
normalized_gpu_type = gpu_type.upper()
|
||||
resolved_git_repo = _resolve_git_repo(git_repo)
|
||||
resolved_git_commit = _resolve_git_commit(git_commit)
|
||||
resolved_pr_number = _resolve_pull_request(pr_number)
|
||||
runner = _select_runner(normalized_gpu_type, num_gpus)
|
||||
|
||||
print(f"Launching {normalized_gpu_type}:{num_gpus} job")
|
||||
print(f"Command: {command}")
|
||||
print(f"Repo: {resolved_git_repo}")
|
||||
print(f"Commit: {resolved_git_commit}")
|
||||
if resolved_pr_number and resolved_pr_number != "false":
|
||||
print(f"PR ref: {resolved_pr_number}")
|
||||
local_patch_b64 = ""
|
||||
if apply_local_patch:
|
||||
patch = _build_local_patch(patch_paths)
|
||||
local_patch_b64 = base64.b64encode(patch.encode("utf-8")).decode("ascii")
|
||||
print(f"Local patch payload: {len(patch)} bytes")
|
||||
|
||||
kwargs = dict(
|
||||
command=command,
|
||||
git_repo=resolved_git_repo,
|
||||
git_commit=resolved_git_commit,
|
||||
pr_number=resolved_pr_number,
|
||||
install_extra=install_extra,
|
||||
build_kernel=build_kernel,
|
||||
env_vars=env_vars,
|
||||
local_patch_b64=local_patch_b64,
|
||||
commit_volume=commit_volume,
|
||||
)
|
||||
if wait:
|
||||
result = runner.remote(**kwargs)
|
||||
print(f"Completed {normalized_gpu_type} job: {result}")
|
||||
return
|
||||
|
||||
function_call = runner.spawn(**kwargs)
|
||||
print(f"Spawned Modal FunctionCall: {function_call.object_id}")
|
||||
print("Poll later with:")
|
||||
print(f" python -c \"import modal; print(modal.FunctionCall.from_id('{function_call.object_id}').get())\"")
|
||||
@@ -0,0 +1,114 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Latent-slice regression tests for Flux2 text-to-image variants.
|
||||
|
||||
Flux2 currently has local parity coverage against the official/reference
|
||||
pipeline, but CI needs a small deterministic regression gate for seeded HF
|
||||
artefacts. Pixel-space comparisons are unnecessarily brittle for this first
|
||||
slot, so the test follows the latent helper pattern used by LTX-2: generate a
|
||||
single-image latent with the production recipe, persist the generated latent,
|
||||
and compare a stable latent signature plus the full tensor against the device
|
||||
reference.
|
||||
|
||||
The default and full-quality parameter maps intentionally carry the same
|
||||
recipe values for now. The ``--ssim-full-quality`` flag still switches the
|
||||
reference tier through ``conftest.py``; separate full-quality recipes can be
|
||||
introduced after the initial Flux2 references have a stable CI window.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
)
|
||||
from fastvideo.tests.ssim.latent_similarity_utils import (
|
||||
run_text_to_latent_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
device_reference_folder = resolve_inference_device_reference_folder(logger)
|
||||
|
||||
FLUX2_MODEL_TO_PARAMS: dict[str, dict[str, object]] = {
|
||||
"black-forest-labs/FLUX.2-klein-4B": {
|
||||
"num_gpus": 1,
|
||||
"model_path": "black-forest-labs/FLUX.2-klein-4B",
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 0,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"fps": 1,
|
||||
},
|
||||
"black-forest-labs/FLUX.2-klein-9B": {
|
||||
"num_gpus": 1,
|
||||
"model_path": "black-forest-labs/FLUX.2-klein-9B",
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 0,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"fps": 1,
|
||||
},
|
||||
}
|
||||
|
||||
FLUX2_FULL_QUALITY_MODEL_TO_PARAMS: dict[str, dict[str, object]] = {
|
||||
model_id: dict(params)
|
||||
for model_id, params in FLUX2_MODEL_TO_PARAMS.items()
|
||||
}
|
||||
|
||||
TEST_PROMPTS: dict[str, str] = {
|
||||
"black-forest-labs/FLUX.2-klein-4B": "a brushed steel espresso machine on a marble counter, morning window light",
|
||||
"black-forest-labs/FLUX.2-klein-9B": "a brushed steel espresso machine on a marble counter, morning window light",
|
||||
}
|
||||
|
||||
SLICE_COSINE_THRESHOLD = 0.96
|
||||
FULL_COSINE_THRESHOLD = 0.99
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="Flux2 SSIM test requires CUDA",
|
||||
)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(FLUX2_MODEL_TO_PARAMS.keys()))
|
||||
def test_flux2_similarity(
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
_ = run_text_to_latent_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=TEST_PROMPTS[model_id],
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=FLUX2_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FLUX2_FULL_QUALITY_MODEL_TO_PARAMS,
|
||||
slice_cosine_threshold=SLICE_COSINE_THRESHOLD,
|
||||
full_cosine_threshold=FULL_COSINE_THRESHOLD,
|
||||
init_kwargs_override={
|
||||
"workload_type": "t2i",
|
||||
"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": (
|
||||
"Flux2KleinPipeline" if "klein" in model_id.lower() else "Flux2Pipeline"
|
||||
),
|
||||
},
|
||||
)
|
||||
@@ -362,7 +362,6 @@ class ValidationCallback(Callback):
|
||||
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=self.validation_random_generator,
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -401,6 +401,7 @@ class WanModel(ModelBase):
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=tc.vsa_sparsity,
|
||||
device=self.device,
|
||||
cache_tile_buf=False,
|
||||
)
|
||||
elif (envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN"):
|
||||
if (not is_vmoba_available() or VideoMobaAttentionMetadataBuilder is None):
|
||||
|
||||
@@ -245,14 +245,16 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
self.generator_ema_2: EMA_FSDP | None = None
|
||||
if (self.training_args.ema_decay is not None) and (self.training_args.ema_decay > 0.0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA with decay=%s", self.training_args.ema_decay)
|
||||
|
||||
# Initialize EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None:
|
||||
self.generator_ema_2 = EMA_FSDP(self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA_2 with decay=%s", self.training_args.ema_decay)
|
||||
ema_enabled = (self.training_args.ema_decay is not None) and (self.training_args.ema_decay > 0.0)
|
||||
if ema_enabled and (self.training_args.ema_start_step <= 0):
|
||||
# Only eager-construct from the cold init weights when averaging starts at step 0.
|
||||
self._build_generator_emas(context="eager init, ema_start_step<=0")
|
||||
elif ema_enabled:
|
||||
# Defer construction to the lazy block in the train loop, which builds the EMA AT
|
||||
# ema_start_step from the already-trained weights. Eager-constructing here would anchor
|
||||
# the shadow to the cold init and leave it base-contaminated (blurry) on short runs.
|
||||
logger.info("Generator EMA deferred: built lazily at ema_start_step=%s from trained weights",
|
||||
self.training_args.ema_start_step)
|
||||
else:
|
||||
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
|
||||
|
||||
@@ -326,22 +328,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
return model
|
||||
return model
|
||||
|
||||
def get_ema_model_copy(self) -> torch.nn.Module | None:
|
||||
"""Get a copy of the model with EMA weights applied."""
|
||||
if self.generator_ema is not None:
|
||||
ema_model = copy.deepcopy(self.transformer)
|
||||
self.generator_ema.copy_to_unwrapped(ema_model)
|
||||
return ema_model
|
||||
return None
|
||||
|
||||
def get_ema_2_model_copy(self) -> torch.nn.Module | None:
|
||||
"""Get a copy of the transformer_2 model with EMA weights applied."""
|
||||
if self.generator_ema_2 is not None and self.transformer_2 is not None:
|
||||
ema_2_model = copy.deepcopy(self.transformer_2)
|
||||
self.generator_ema_2.copy_to_unwrapped(ema_2_model)
|
||||
return ema_2_model
|
||||
return None
|
||||
|
||||
def is_ema_ready(self, current_step: int | None = None):
|
||||
"""Check if EMA is ready for use (after ema_start_step)."""
|
||||
if current_step is None:
|
||||
@@ -361,68 +347,61 @@ class DistillationPipeline(TrainingPipeline):
|
||||
try:
|
||||
# Save main transformer EMA
|
||||
if self.generator_ema is not None:
|
||||
ema_model = self.get_ema_model_copy()
|
||||
if ema_model is None:
|
||||
logger.warning("Failed to create EMA model copy")
|
||||
else:
|
||||
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from fastvideo.training.training_utils import (custom_to_hf_state_dict,
|
||||
gather_state_dict_on_cpu_rank0)
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
|
||||
from fastvideo.training.training_utils import (custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
# Swap EMA weights into the live FSDP module in place (no deepcopy) and gather the
|
||||
# full state dict within the context; weights are restored on exit.
|
||||
with self.generator_ema.apply_to_model(self.transformer):
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(self.transformer, device=None)
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(cpu_state, ema_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(cpu_state,
|
||||
self.transformer.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
|
||||
config_dict = ema_model.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
# deepcopy so deleting "dtype" doesn't mutate the live model's hf_config
|
||||
config_dict = copy.deepcopy(self.transformer.hf_config)
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
|
||||
del ema_model
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
|
||||
# Save transformer_2 EMA
|
||||
if self.generator_ema_2 is not None:
|
||||
ema_2_model = self.get_ema_2_model_copy()
|
||||
if ema_2_model is None:
|
||||
logger.warning("Failed to create EMA_2 model copy")
|
||||
else:
|
||||
ema_2_save_dir = os.path.join(output_dir, f"ema_2_checkpoint-{step}")
|
||||
os.makedirs(ema_2_save_dir, exist_ok=True)
|
||||
if self.generator_ema_2 is not None and self.transformer_2 is not None:
|
||||
ema_2_save_dir = os.path.join(output_dir, f"ema_2_checkpoint-{step}")
|
||||
os.makedirs(ema_2_save_dir, exist_ok=True)
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from fastvideo.training.training_utils import (custom_to_hf_state_dict,
|
||||
gather_state_dict_on_cpu_rank0)
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(ema_2_model, device=None)
|
||||
from fastvideo.training.training_utils import (custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
with self.generator_ema_2.apply_to_model(self.transformer_2):
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(self.transformer_2, device=None)
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path_2 = os.path.join(ema_2_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(cpu_state_2,
|
||||
ema_2_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
if self.global_rank == 0:
|
||||
weight_path_2 = os.path.join(ema_2_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(cpu_state_2,
|
||||
self.transformer_2.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
|
||||
config_dict_2 = ema_2_model.hf_config
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"]
|
||||
config_path_2 = os.path.join(ema_2_save_dir, "config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
# deepcopy so deleting "dtype" doesn't mutate the live model's hf_config
|
||||
config_dict_2 = copy.deepcopy(self.transformer_2.hf_config)
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"]
|
||||
config_path_2 = os.path.join(ema_2_save_dir, "config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
|
||||
logger.info("EMA_2 weights saved to %s", weight_path_2)
|
||||
|
||||
del ema_2_model
|
||||
logger.info("EMA_2 weights saved to %s", weight_path_2)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to save EMA weights: %s", str(e))
|
||||
@@ -919,11 +898,45 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
|
||||
return training_batch
|
||||
|
||||
def _build_generator_emas(self, context: str = "") -> None:
|
||||
# Idempotently construct whichever generator EMA shadows are missing, per-expert and
|
||||
# decoupled. Safe to call repeatedly; no-op once both exist or when EMA is disabled.
|
||||
if (self.training_args.ema_decay is None) or (self.training_args.ema_decay <= 0.0):
|
||||
return
|
||||
suffix = f" [{context}]" if context else ""
|
||||
if self.generator_ema is None:
|
||||
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Built generator EMA (decay=%s)%s", self.training_args.ema_decay, suffix)
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Built generator EMA_2 (decay=%s)%s", self.training_args.ema_decay, suffix)
|
||||
|
||||
def _build_deferred_ema_for_resume(self) -> None:
|
||||
# Build a deferred EMA before checkpoint load only when a saved shard exists, so the
|
||||
# shadow reloads instead of being skipped and rebuilt fresh. Gating on shard existence
|
||||
# matters: building when none exists would reintroduce cold-init contamination.
|
||||
if (self.training_args.ema_decay is None) or (self.training_args.ema_decay <= 0.0):
|
||||
return
|
||||
|
||||
ema_shard_dir = os.path.join(self.training_args.resume_from_checkpoint, "ema_local_shard")
|
||||
|
||||
if (self.generator_ema is None
|
||||
and os.path.exists(os.path.join(ema_shard_dir, f"generator_ema_rank{self.global_rank}.pt"))):
|
||||
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Pre-built generator EMA for resume from existing shard")
|
||||
|
||||
if (self.transformer_2 is not None and self.generator_ema_2 is None
|
||||
and os.path.exists(os.path.join(ema_shard_dir, f"generator_ema_2_rank{self.global_rank}.pt"))):
|
||||
self.generator_ema_2 = EMA_FSDP(self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Pre-built generator EMA_2 for resume from existing shard")
|
||||
|
||||
def _resume_from_checkpoint(self) -> None:
|
||||
"""Resume training from checkpoint with distillation models."""
|
||||
|
||||
logger.info("Loading distillation checkpoint from %s", self.training_args.resume_from_checkpoint)
|
||||
|
||||
self._build_deferred_ema_for_resume()
|
||||
|
||||
resumed_step = load_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
@@ -1332,15 +1345,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
if (step >= self.training_args.ema_start_step) and \
|
||||
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Created generator EMA at step %s with decay=%s", step, self.training_args.ema_decay)
|
||||
|
||||
# Create EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Created generator EMA_2 at step %s with decay=%s", step, self.training_args.ema_decay)
|
||||
if step >= self.training_args.ema_start_step:
|
||||
self._build_generator_emas(context=f"lazy @ step {step}")
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -405,7 +405,6 @@ class MatrixGame2ARDiffusionPipeline(TrainingPipeline):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -351,7 +351,6 @@ class MatrixGame2ODEInitTrainingPipeline(TrainingPipeline):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -853,7 +853,6 @@ class MatrixGame2SelfForcingDistillationPipeline(SelfForcingDistillationPipeline
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -172,7 +172,6 @@ class MatrixGame2TrainingPipeline(TrainingPipeline):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -852,15 +852,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
if (step >= self.training_args.ema_start_step) and \
|
||||
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info("Created generator EMA at step %s with decay=%s", step, self.training_args.ema_decay)
|
||||
|
||||
# Create EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Created generator EMA_2 at step %s with decay=%s", step, self.training_args.ema_decay)
|
||||
if step >= self.training_args.ema_start_step:
|
||||
self._build_generator_emas(context=f"lazy @ step {step}")
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -359,7 +359,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_timestep=training_batch.timesteps,
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=current_vsa_sparsity,
|
||||
device=get_local_torch_device())
|
||||
device=get_local_torch_device(),
|
||||
cache_tile_buf=False)
|
||||
elif envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
if not vmoba_available:
|
||||
raise ImportError("FASTVIDEO_ATTENTION_BACKEND is set to VMOBA_ATTN, "
|
||||
@@ -690,7 +691,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=self.validation_random_generator,
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -613,18 +613,10 @@ def load_distillation_checkpoint(
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA separately if saved in rank0_full mode
|
||||
# Load EMA separately
|
||||
if generator_ema is not None:
|
||||
try:
|
||||
if getattr(generator_ema, "mode", None) == "rank0_full":
|
||||
ema_path = os.path.join(checkpoint_path, "ema", "generator_ema.pt")
|
||||
if rank == 0 and os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info("rank: %s, generator EMA (rank0_full) loaded from %s", rank, ema_path)
|
||||
elif rank == 0:
|
||||
logger.info("rank: %s, generator EMA file not found at %s; skipping", rank, ema_path)
|
||||
elif getattr(generator_ema, "mode", None) == "local_shard":
|
||||
if getattr(generator_ema, "mode", None) == "local_shard":
|
||||
ema_path = os.path.join(checkpoint_path, "ema_local_shard", f"generator_ema_rank{rank}.pt")
|
||||
if os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
@@ -632,9 +624,28 @@ def load_distillation_checkpoint(
|
||||
logger.info("rank: %s, generator EMA shard (local_shard) loaded from %s", rank, ema_path)
|
||||
else:
|
||||
logger.info("rank: %s, generator EMA shard file not found at %s; skipping", rank, ema_path)
|
||||
else:
|
||||
logger.info("rank: %s, generator EMA mode %s not supported for resume; skipping", rank,
|
||||
getattr(generator_ema, "mode", None))
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load generator EMA: %s", rank, str(e))
|
||||
|
||||
# Load EMA_2 from its shard (symmetric with generator_ema above)
|
||||
if generator_ema_2 is not None:
|
||||
try:
|
||||
if getattr(generator_ema_2, "mode", None) == "local_shard":
|
||||
ema_2_path = os.path.join(checkpoint_path, "ema_local_shard", f"generator_ema_2_rank{rank}.pt")
|
||||
if os.path.exists(ema_2_path):
|
||||
generator_ema_2.load_state_dict(torch.load(ema_2_path, map_location="cpu"))
|
||||
logger.info("rank: %s, generator EMA_2 shard (local_shard) loaded from %s", rank, ema_2_path)
|
||||
else:
|
||||
logger.info("rank: %s, generator EMA_2 shard file not found at %s; skipping", rank, ema_2_path)
|
||||
else:
|
||||
logger.info("rank: %s, generator EMA_2 mode %s not supported for resume; skipping", rank,
|
||||
getattr(generator_ema_2, "mode", None))
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load generator EMA_2: %s", rank, str(e))
|
||||
|
||||
# Load generator_2 distributed checkpoint (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
generator_2_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint", "generator_2")
|
||||
@@ -665,18 +676,6 @@ def load_distillation_checkpoint(
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA_2 state if available and generator_ema_2 is provided
|
||||
if generator_ema_2 is not None:
|
||||
try:
|
||||
ema_2_state = generator_2_states.get("ema")
|
||||
if ema_2_state is not None:
|
||||
generator_ema_2.load_state_dict(ema_2_state)
|
||||
logger.info("rank: %s, generator_2 EMA state loaded successfully", rank)
|
||||
else:
|
||||
logger.info("rank: %s, no EMA_2 state found in checkpoint", rank)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA_2 state: %s", rank, str(e))
|
||||
else:
|
||||
logger.info("rank: %s, generator_2 checkpoint not found, skipping", rank)
|
||||
|
||||
|
||||
@@ -114,7 +114,6 @@ class WanI2VDistillationPipeline(DistillationPipeline):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -157,7 +157,6 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.1.7"
|
||||
__version__ = "0.2.0"
|
||||
|
||||
@@ -138,7 +138,9 @@ class MultiprocExecutor(Executor):
|
||||
result_batch = ForwardBatch(data_type=forward_batch.data_type,
|
||||
output=output,
|
||||
logging_info=logging_info,
|
||||
extra=extra)
|
||||
extra=extra,
|
||||
trajectory_latents=responses[0].get("trajectory_latents"),
|
||||
trajectory_timesteps=responses[0].get("trajectory_timesteps"))
|
||||
|
||||
return result_batch
|
||||
|
||||
@@ -699,6 +701,8 @@ class WorkerMultiprocProc:
|
||||
"output_batch": result,
|
||||
"logging_info": logging_info,
|
||||
"extra": extra,
|
||||
"trajectory_latents": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(method, *args, **kwargs)
|
||||
|
||||
@@ -307,6 +307,8 @@ class RayDistributedExecutor(Executor):
|
||||
data_type=forward_batch.data_type,
|
||||
output=output,
|
||||
logging_info=logging_info,
|
||||
trajectory_latents=responses[0].trajectory_latents,
|
||||
trajectory_timesteps=responses[0].trajectory_timesteps,
|
||||
)
|
||||
return result_batch
|
||||
|
||||
|
||||
+11
-4
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.7"
|
||||
version = "0.2.0"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
@@ -26,7 +26,7 @@ dependencies = [
|
||||
"sentencepiece>=0.2.0",
|
||||
"timm>=1.0.11",
|
||||
"peft>=0.15.0",
|
||||
"diffusers>=0.33.1",
|
||||
"diffusers>=0.38.0",
|
||||
"torch==2.11.0",
|
||||
"torchvision",
|
||||
"torchaudio",
|
||||
@@ -106,6 +106,7 @@ torchaudio = [
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
|
||||
]
|
||||
imagebind = { git = "https://github.com/facebookresearch/ImageBind.git", rev = "53680b02d7e37b19b124fa37bae4b6c98c38f5be" }
|
||||
flash-attn-cute = { git = "https://github.com/XOR-op/flash-attention.git", branch = "fa4-compile", subdirectory = "flash_attn/cute" }
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cpu"
|
||||
@@ -142,10 +143,13 @@ test = [
|
||||
# separately as documented in the README. decord has no aarch64 wheels
|
||||
# (upstream effectively unmaintained); it's split into its own extra so
|
||||
# the main eval extras install on ARM systems too. eval/io/video.py has
|
||||
# a PyAV fallback that covers the video-decode path without decord.
|
||||
# a PyAV fallback that covers the video-decode path without decord. The
|
||||
# judge.* metrics (eval-judge) call a remote API and need a key, so they
|
||||
# are opt-in and intentionally kept out of the default [eval] set.
|
||||
eval-vbench = ["openai-clip", "pyiqa", "easydict"]
|
||||
eval-physics-iq = []
|
||||
eval-audio = ["jiwer", "librosa", "pyloudnorm", "hear21passt", "audiobox_aesthetics", "imagebind", "pytorchvideo"]
|
||||
eval-judge = ["google-genai"]
|
||||
eval = [
|
||||
"lpips", "ptlflow", "qwen-vl-utils",
|
||||
"fastvideo[eval-vbench]", "fastvideo[eval-physics-iq]",
|
||||
@@ -172,7 +176,10 @@ streaming = [
|
||||
dreamverse = [
|
||||
"uvicorn[standard]>=0.41.0",
|
||||
"cerebras-cloud-sdk",
|
||||
"flash-attn-cute @ git+https://github.com/XOR-op/flash-attention.git@fa4-compile#subdirectory=flash_attn/cute",
|
||||
# PyPI forbids direct URL deps in published metadata; pin the fork via
|
||||
# [tool.uv.sources] above (same pattern as imagebind) so the wheel stays
|
||||
# publishable while uv workspace installs still get the fork.
|
||||
"flash-attn-cute",
|
||||
"flashinfer-python",
|
||||
"openai>=1.40",
|
||||
]
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.1.7"
|
||||
version = "0.2.0"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -0,0 +1,527 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# pyright: reportAny=false, reportExplicitAny=false, reportUnknownMemberType=false
|
||||
# pyright: reportUnknownVariableType=false, reportUnusedCallResult=false
|
||||
"""Convert FLUX.2 Klein weights into a FastVideo-loadable layout.
|
||||
|
||||
The published ``black-forest-labs/FLUX.2-klein-4B`` repo contains two useful
|
||||
transformer surfaces:
|
||||
|
||||
* ``flux-2-klein-4b.safetensors``: BFL's compact raw transformer checkpoint.
|
||||
Its double-stream attention projections are fused as ``img_attn.qkv`` and
|
||||
``txt_attn.qkv``.
|
||||
* ``transformer/diffusion_pytorch_model.safetensors``: Diffusers/FastVideo
|
||||
names. These keys already match ``Flux2Transformer2DModel.state_dict()``.
|
||||
|
||||
By default this script prefers the raw root checkpoint when present, converts it
|
||||
to the FastVideo native transformer key surface, reloads/resaves the VAE, copies
|
||||
the HF-backed Qwen3 text encoder/tokenizer and scheduler, and emits a standard
|
||||
Diffusers-style FastVideo repo:
|
||||
|
||||
<dst>/
|
||||
model_index.json
|
||||
transformer/{config.json,diffusion_pytorch_model*.safetensors}
|
||||
vae/{config.json,diffusion_pytorch_model*.safetensors}
|
||||
text_encoder/...
|
||||
tokenizer/...
|
||||
scheduler/scheduler_config.json
|
||||
|
||||
Qwen3 and Mistral3 are intentionally copied as Transformers passthrough
|
||||
components in the standard layout: ``qwen3.py`` and ``mistral3.py`` both route
|
||||
Flux2 text encoding through ``from_pretrained_local()`` for exact HF parity.
|
||||
|
||||
Example:
|
||||
python scripts/checkpoint_conversion/convert_flux2_klein.py \
|
||||
--src black-forest-labs/FLUX.2-klein-4B \
|
||||
--dst converted_weights/flux2-klein-4b-fastvideo \
|
||||
--overwrite
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
|
||||
try:
|
||||
from huggingface_hub import save_torch_state_dict, snapshot_download
|
||||
except ImportError:
|
||||
save_torch_state_dict = None
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
DEFAULT_REPO_ID = "black-forest-labs/FLUX.2-klein-4B"
|
||||
RAW_TRANSFORMER_FILENAME = "flux-2-klein-4b.safetensors"
|
||||
DIFFUSION_WEIGHTS_BASENAME = "diffusion_pytorch_model"
|
||||
|
||||
BASE_SNAPSHOT_ALLOW_PATTERNS = (
|
||||
"model_index.json",
|
||||
"transformer/config.json",
|
||||
"vae/config.json",
|
||||
"vae/*.safetensors",
|
||||
"vae/*.safetensors.index.json",
|
||||
"text_encoder/*",
|
||||
"tokenizer/*",
|
||||
"scheduler/*",
|
||||
)
|
||||
|
||||
PASSTHROUGH_SUBFOLDERS = ("text_encoder", "tokenizer", "scheduler")
|
||||
|
||||
DEFAULT_MODEL_INDEX: dict[str, Any] = {
|
||||
"_class_name": "Flux2KleinPipeline",
|
||||
"_diffusers_version": "0.37.0.dev0",
|
||||
"is_distilled": True,
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "Qwen3ForCausalLM"],
|
||||
"tokenizer": ["transformers", "Qwen2TokenizerFast"],
|
||||
"transformer": ["diffusers", "Flux2Transformer2DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLFlux2"],
|
||||
}
|
||||
|
||||
DEFAULT_SCHEDULER_CONFIG: dict[str, Any] = {
|
||||
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.37.0.dev0",
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"invert_sigmas": False,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"shift_terminal": None,
|
||||
"stochastic_sampling": False,
|
||||
"time_shift_type": "exponential",
|
||||
"use_beta_sigmas": False,
|
||||
"use_dynamic_shifting": True,
|
||||
"use_exponential_sigmas": False,
|
||||
"use_karras_sigmas": False,
|
||||
}
|
||||
|
||||
TRANSFORMER_REQUIRED_KEYS = (
|
||||
"x_embedder.weight",
|
||||
"context_embedder.weight",
|
||||
"time_guidance_embed.timestep_embedder.linear_1.weight",
|
||||
"time_guidance_embed.timestep_embedder.linear_2.weight",
|
||||
"double_stream_modulation_img.linear.weight",
|
||||
"double_stream_modulation_txt.linear.weight",
|
||||
"single_stream_modulation.linear.weight",
|
||||
"norm_out.linear.weight",
|
||||
"proj_out.weight",
|
||||
)
|
||||
|
||||
VAE_REQUIRED_KEYS = (
|
||||
"encoder.conv_in.weight",
|
||||
"decoder.conv_out.weight",
|
||||
"quant_conv.weight",
|
||||
"post_quant_conv.weight",
|
||||
"bn.running_mean",
|
||||
"bn.running_var",
|
||||
)
|
||||
|
||||
|
||||
class ConversionError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def _resolve_hf_token() -> str | bool | None:
|
||||
try:
|
||||
from fastvideo.utils import resolve_hf_token
|
||||
|
||||
token = resolve_hf_token()
|
||||
return token if token else None
|
||||
except Exception:
|
||||
return os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN") or None
|
||||
|
||||
|
||||
def _snapshot_allow_patterns(transformer_source: str) -> list[str]:
|
||||
patterns: list[str] = list(BASE_SNAPSHOT_ALLOW_PATTERNS)
|
||||
if transformer_source == "diffusers":
|
||||
patterns.extend(("transformer/*.safetensors", "transformer/*.safetensors.index.json"))
|
||||
else:
|
||||
patterns.append(RAW_TRANSFORMER_FILENAME)
|
||||
return patterns
|
||||
|
||||
|
||||
def _resolve_src(src: str, revision: str | None, cache_dir: str | None, transformer_source: str) -> Path:
|
||||
local = Path(src).expanduser()
|
||||
if local.exists():
|
||||
return local
|
||||
|
||||
if snapshot_download is None:
|
||||
raise ConversionError("huggingface_hub is required when --src is a repo id")
|
||||
|
||||
return Path(
|
||||
snapshot_download(
|
||||
repo_id=src,
|
||||
revision=revision,
|
||||
cache_dir=cache_dir,
|
||||
token=_resolve_hf_token(),
|
||||
allow_patterns=_snapshot_allow_patterns(transformer_source),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _read_json(path: Path) -> dict[str, Any]:
|
||||
with path.open("r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _write_json(path: Path, payload: dict[str, Any]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with path.open("w", encoding="utf-8") as f:
|
||||
json.dump(payload, f, indent=2)
|
||||
f.write("\n")
|
||||
|
||||
|
||||
def _prepare_output_dir(dst: Path, overwrite: bool) -> None:
|
||||
if dst.exists() and any(dst.iterdir()):
|
||||
if not overwrite:
|
||||
raise FileExistsError(f"Output directory is not empty: {dst}. Pass --overwrite to replace it.")
|
||||
shutil.rmtree(dst)
|
||||
dst.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
def _find_safetensors_files(component_dir: Path, basename: str = DIFFUSION_WEIGHTS_BASENAME) -> list[Path]:
|
||||
if component_dir.is_file():
|
||||
return [component_dir]
|
||||
|
||||
index_path = component_dir / f"{basename}.safetensors.index.json"
|
||||
if index_path.exists():
|
||||
index = _read_json(index_path)
|
||||
return sorted({component_dir / shard for shard in index["weight_map"].values()})
|
||||
|
||||
single = component_dir / f"{basename}.safetensors"
|
||||
if single.exists():
|
||||
return [single]
|
||||
|
||||
return sorted(component_dir.glob("*.safetensors"))
|
||||
|
||||
|
||||
def _load_safetensors(files: list[Path]) -> OrderedDict[str, torch.Tensor]:
|
||||
if not files:
|
||||
raise FileNotFoundError("No safetensors files found")
|
||||
|
||||
state: OrderedDict[str, torch.Tensor] = OrderedDict()
|
||||
for file in files:
|
||||
with safe_open(str(file), framework="pt", device="cpu") as handle:
|
||||
for key in handle.keys():
|
||||
if key in state:
|
||||
raise ConversionError(f"Duplicate tensor key {key!r} while reading {file}")
|
||||
state[key] = handle.get_tensor(key)
|
||||
return state
|
||||
|
||||
|
||||
def _write_state_dict(
|
||||
state: OrderedDict[str, torch.Tensor],
|
||||
output_dir: Path,
|
||||
max_shard_size: str,
|
||||
) -> None:
|
||||
if save_torch_state_dict is None:
|
||||
raise ConversionError("huggingface_hub.save_torch_state_dict is required to write sharded safetensors")
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_torch_state_dict(
|
||||
state,
|
||||
output_dir,
|
||||
filename_pattern=f"{DIFFUSION_WEIGHTS_BASENAME}" + "{suffix}.safetensors",
|
||||
max_shard_size=max_shard_size,
|
||||
metadata={"format": "pt"},
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
|
||||
def _split_qkv(weight: torch.Tensor, source_key: str) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
if weight.shape[0] % 3 != 0:
|
||||
raise ConversionError(f"Expected first dim divisible by 3 for {source_key}, got {tuple(weight.shape)}")
|
||||
q_weight, k_weight, v_weight = torch.chunk(weight, 3, dim=0)
|
||||
return q_weight, k_weight, v_weight
|
||||
|
||||
|
||||
def _convert_raw_transformer_key(
|
||||
key: str,
|
||||
value: torch.Tensor,
|
||||
output: OrderedDict[str, torch.Tensor],
|
||||
) -> bool:
|
||||
literal_map = {
|
||||
"img_in.weight": "x_embedder.weight",
|
||||
"txt_in.weight": "context_embedder.weight",
|
||||
"time_in.in_layer.weight": "time_guidance_embed.timestep_embedder.linear_1.weight",
|
||||
"time_in.out_layer.weight": "time_guidance_embed.timestep_embedder.linear_2.weight",
|
||||
"double_stream_modulation_img.lin.weight": "double_stream_modulation_img.linear.weight",
|
||||
"double_stream_modulation_txt.lin.weight": "double_stream_modulation_txt.linear.weight",
|
||||
"single_stream_modulation.lin.weight": "single_stream_modulation.linear.weight",
|
||||
"final_layer.adaLN_modulation.1.weight": "norm_out.linear.weight",
|
||||
"final_layer.linear.weight": "proj_out.weight",
|
||||
}
|
||||
if key in literal_map:
|
||||
output[literal_map[key]] = value
|
||||
return True
|
||||
|
||||
match = re.fullmatch(r"double_blocks\.(\d+)\.(img|txt)_attn\.qkv\.weight", key)
|
||||
if match:
|
||||
block, stream = match.groups()
|
||||
q_weight, k_weight, v_weight = _split_qkv(value, key)
|
||||
if stream == "img":
|
||||
prefix = f"transformer_blocks.{block}.attn"
|
||||
output[f"{prefix}.to_q.weight"] = q_weight
|
||||
output[f"{prefix}.to_k.weight"] = k_weight
|
||||
output[f"{prefix}.to_v.weight"] = v_weight
|
||||
else:
|
||||
prefix = f"transformer_blocks.{block}.attn"
|
||||
output[f"{prefix}.add_q_proj.weight"] = q_weight
|
||||
output[f"{prefix}.add_k_proj.weight"] = k_weight
|
||||
output[f"{prefix}.add_v_proj.weight"] = v_weight
|
||||
return True
|
||||
|
||||
double_rewrites = (
|
||||
(r"double_blocks\.(\d+)\.img_attn\.proj\.weight", r"transformer_blocks.\1.attn.to_out.0.weight"),
|
||||
(r"double_blocks\.(\d+)\.txt_attn\.proj\.weight", r"transformer_blocks.\1.attn.to_add_out.weight"),
|
||||
(r"double_blocks\.(\d+)\.img_attn\.norm\.query_norm\.scale", r"transformer_blocks.\1.attn.norm_q.weight"),
|
||||
(r"double_blocks\.(\d+)\.img_attn\.norm\.key_norm\.scale", r"transformer_blocks.\1.attn.norm_k.weight"),
|
||||
(r"double_blocks\.(\d+)\.txt_attn\.norm\.query_norm\.scale", r"transformer_blocks.\1.attn.norm_added_q.weight"),
|
||||
(r"double_blocks\.(\d+)\.txt_attn\.norm\.key_norm\.scale", r"transformer_blocks.\1.attn.norm_added_k.weight"),
|
||||
(r"double_blocks\.(\d+)\.img_mlp\.0\.weight", r"transformer_blocks.\1.ff.linear_in.weight"),
|
||||
(r"double_blocks\.(\d+)\.img_mlp\.2\.weight", r"transformer_blocks.\1.ff.linear_out.weight"),
|
||||
(r"double_blocks\.(\d+)\.txt_mlp\.0\.weight", r"transformer_blocks.\1.ff_context.linear_in.weight"),
|
||||
(r"double_blocks\.(\d+)\.txt_mlp\.2\.weight", r"transformer_blocks.\1.ff_context.linear_out.weight"),
|
||||
(r"single_blocks\.(\d+)\.linear1\.weight", r"single_transformer_blocks.\1.attn.to_qkv_mlp_proj.weight"),
|
||||
(r"single_blocks\.(\d+)\.linear2\.weight", r"single_transformer_blocks.\1.attn.to_out.weight"),
|
||||
(r"single_blocks\.(\d+)\.norm\.query_norm\.scale", r"single_transformer_blocks.\1.attn.norm_q.weight"),
|
||||
(r"single_blocks\.(\d+)\.norm\.key_norm\.scale", r"single_transformer_blocks.\1.attn.norm_k.weight"),
|
||||
)
|
||||
for pattern, replacement in double_rewrites:
|
||||
if re.fullmatch(pattern, key):
|
||||
output[re.sub(pattern, replacement, key)] = value
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def convert_raw_transformer(raw_state: OrderedDict[str, torch.Tensor]) -> OrderedDict[str, torch.Tensor]:
|
||||
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
|
||||
unexpected: list[str] = []
|
||||
for key, value in raw_state.items():
|
||||
if not _convert_raw_transformer_key(key, value, converted):
|
||||
unexpected.append(key)
|
||||
|
||||
if unexpected:
|
||||
raise ConversionError(
|
||||
"Unmapped raw Flux2 Klein transformer keys:\n" + "\n".join(f" - {key}" for key in unexpected)
|
||||
)
|
||||
return converted
|
||||
|
||||
|
||||
def convert_diffusers_transformer(state: OrderedDict[str, torch.Tensor]) -> OrderedDict[str, torch.Tensor]:
|
||||
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
|
||||
for key, value in state.items():
|
||||
if key.startswith("transformer."):
|
||||
key = key[len("transformer."):]
|
||||
converted[key] = value
|
||||
return converted
|
||||
|
||||
|
||||
def convert_vae(state: OrderedDict[str, torch.Tensor]) -> OrderedDict[str, torch.Tensor]:
|
||||
return OrderedDict(state.items())
|
||||
|
||||
|
||||
def _infer_block_count(keys: set[str], prefix: str) -> int:
|
||||
count = 0
|
||||
for key in keys:
|
||||
if key.startswith(prefix):
|
||||
suffix = key[len(prefix):]
|
||||
first = suffix.split(".", 1)[0]
|
||||
if first.isdigit():
|
||||
count = max(count, int(first) + 1)
|
||||
return count
|
||||
|
||||
|
||||
def _validate_required_keys(component: str, state: OrderedDict[str, torch.Tensor], required: tuple[str, ...]) -> None:
|
||||
missing = [key for key in required if key not in state]
|
||||
if missing:
|
||||
raise ConversionError(f"{component} conversion missing required keys: {missing}")
|
||||
|
||||
|
||||
def _validate_transformer(state: OrderedDict[str, torch.Tensor]) -> None:
|
||||
_validate_required_keys("transformer", state, TRANSFORMER_REQUIRED_KEYS)
|
||||
raw_leftovers = [key for key in state if key.startswith(("double_blocks.", "single_blocks."))]
|
||||
if raw_leftovers:
|
||||
raise ConversionError(f"Raw transformer keys leaked into output: {raw_leftovers[:8]}")
|
||||
keys = set(state)
|
||||
double_blocks = _infer_block_count(keys, "transformer_blocks.")
|
||||
single_blocks = _infer_block_count(keys, "single_transformer_blocks.")
|
||||
print(
|
||||
" transformer keys: " +
|
||||
f"{len(state)} total, {double_blocks} double blocks, {single_blocks} single blocks"
|
||||
)
|
||||
|
||||
|
||||
def _validate_vae(state: OrderedDict[str, torch.Tensor]) -> None:
|
||||
_validate_required_keys("vae", state, VAE_REQUIRED_KEYS)
|
||||
print(f" vae keys: {len(state)} total")
|
||||
|
||||
|
||||
def _component_config(src_dir: Path, component: str, class_name: str) -> dict[str, Any]:
|
||||
config_path = src_dir / component / "config.json"
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Missing {component} config: {config_path}")
|
||||
config = _read_json(config_path)
|
||||
config["_class_name"] = class_name
|
||||
config.pop("_name_or_path", None)
|
||||
return config
|
||||
|
||||
|
||||
def _copy_passthrough_subfolder(src_dir: Path, dst_dir: Path, subfolder: str) -> None:
|
||||
src = src_dir / subfolder
|
||||
if not src.is_dir():
|
||||
raise FileNotFoundError(f"Missing passthrough subfolder: {src}")
|
||||
dst = dst_dir / subfolder
|
||||
shutil.copytree(src, dst)
|
||||
print(f" copied {subfolder}/")
|
||||
|
||||
|
||||
def _copy_or_write_scheduler(src_dir: Path, dst_dir: Path) -> None:
|
||||
src = src_dir / "scheduler"
|
||||
dst = dst_dir / "scheduler"
|
||||
if src.is_dir():
|
||||
shutil.copytree(src, dst)
|
||||
print(" copied scheduler/")
|
||||
return
|
||||
_write_json(dst / "scheduler_config.json", DEFAULT_SCHEDULER_CONFIG)
|
||||
print(" wrote default scheduler/scheduler_config.json")
|
||||
|
||||
|
||||
def _build_model_index(src_dir: Path, source_label: str) -> dict[str, Any]:
|
||||
index_path = src_dir / "model_index.json"
|
||||
if index_path.exists():
|
||||
index = _read_json(index_path)
|
||||
else:
|
||||
index = dict(DEFAULT_MODEL_INDEX)
|
||||
index.update(
|
||||
{
|
||||
"_class_name": "Flux2KleinPipeline",
|
||||
"is_distilled": True,
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "Qwen3ForCausalLM"],
|
||||
"tokenizer": ["transformers", "Qwen2TokenizerFast"],
|
||||
"transformer": ["diffusers", "Flux2Transformer2DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLFlux2"],
|
||||
"_fastvideo_converted_from": source_label,
|
||||
}
|
||||
)
|
||||
return index
|
||||
|
||||
|
||||
def _select_transformer_source(src_dir: Path, mode: str) -> tuple[str, Path]:
|
||||
if src_dir.is_file():
|
||||
if mode == "diffusers":
|
||||
raise FileNotFoundError(f"Cannot use --transformer-source diffusers with a file source: {src_dir}")
|
||||
return "raw", src_dir
|
||||
|
||||
raw_path = src_dir / RAW_TRANSFORMER_FILENAME
|
||||
diffusers_dir = src_dir / "transformer"
|
||||
|
||||
if mode == "raw":
|
||||
if not raw_path.exists():
|
||||
raise FileNotFoundError(f"Requested raw transformer but missing {raw_path}")
|
||||
return "raw", raw_path
|
||||
if mode == "diffusers":
|
||||
if not diffusers_dir.is_dir():
|
||||
raise FileNotFoundError(f"Requested diffusers transformer but missing {diffusers_dir}")
|
||||
return "diffusers", diffusers_dir
|
||||
if raw_path.exists():
|
||||
return "raw", raw_path
|
||||
if diffusers_dir.is_dir():
|
||||
return "diffusers", diffusers_dir
|
||||
raise FileNotFoundError(f"Missing transformer weights under {src_dir}")
|
||||
|
||||
|
||||
def convert(src: str, dst: str, revision: str | None, cache_dir: str | None, max_shard_size: str,
|
||||
transformer_source: str, overwrite: bool) -> None:
|
||||
resolved_src = _resolve_src(src, revision=revision, cache_dir=cache_dir, transformer_source=transformer_source)
|
||||
src_dir = resolved_src.parent if resolved_src.is_file() else resolved_src
|
||||
dst_dir = Path(dst).expanduser()
|
||||
_prepare_output_dir(dst_dir, overwrite=overwrite)
|
||||
|
||||
print(f"src: {src_dir}")
|
||||
print(f"dst: {dst_dir}")
|
||||
|
||||
print("\n[1/5] Converting transformer:")
|
||||
source_kind, source_path = _select_transformer_source(resolved_src, transformer_source)
|
||||
if source_kind == "raw":
|
||||
print(f" using raw BFL transformer: {source_path.name}")
|
||||
transformer_state = convert_raw_transformer(_load_safetensors([source_path]))
|
||||
else:
|
||||
print(" using Diffusers transformer subfolder")
|
||||
transformer_state = convert_diffusers_transformer(_load_safetensors(_find_safetensors_files(source_path)))
|
||||
_validate_transformer(transformer_state)
|
||||
transformer_dir = dst_dir / "transformer"
|
||||
_write_state_dict(transformer_state, transformer_dir, max_shard_size=max_shard_size)
|
||||
_write_json(transformer_dir / "config.json", _component_config(src_dir, "transformer", "Flux2Transformer2DModel"))
|
||||
|
||||
print("\n[2/5] Converting VAE:")
|
||||
vae_dir = src_dir / "vae"
|
||||
vae_state = convert_vae(_load_safetensors(_find_safetensors_files(vae_dir)))
|
||||
_validate_vae(vae_state)
|
||||
output_vae_dir = dst_dir / "vae"
|
||||
_write_state_dict(vae_state, output_vae_dir, max_shard_size=max_shard_size)
|
||||
_write_json(output_vae_dir / "config.json", _component_config(src_dir, "vae", "AutoencoderKLFlux2"))
|
||||
|
||||
print("\n[3/5] Copying HF-backed encoders/tokenizer/scheduler:")
|
||||
# Current FastVideo Flux2 text encoders intentionally use HF passthrough:
|
||||
# Qwen3ForCausalLM.from_pretrained_local() and Mistral3ForConditionalGeneration.from_pretrained_local().
|
||||
# Rewriting their weights to native fused QKV names would make the standard loader call the wrong HF class.
|
||||
_copy_passthrough_subfolder(src_dir, dst_dir, "text_encoder")
|
||||
_copy_passthrough_subfolder(src_dir, dst_dir, "tokenizer")
|
||||
_copy_or_write_scheduler(src_dir, dst_dir)
|
||||
|
||||
print("\n[4/5] Writing model_index.json:")
|
||||
_write_json(dst_dir / "model_index.json", _build_model_index(src_dir, src))
|
||||
print(" wrote model_index.json")
|
||||
|
||||
print("\n[5/5] Summary:")
|
||||
if source_kind == "raw":
|
||||
print(" transformer mapping: raw BFL -> FastVideo native")
|
||||
else:
|
||||
print(" transformer mapping: Diffusers/FastVideo identity")
|
||||
print(" vae mapping: Diffusers/FastVideo identity")
|
||||
print(" text encoder mapping: HF Qwen3 passthrough (FastVideo loader uses from_pretrained_local)")
|
||||
print(f" max shard size: {max_shard_size}")
|
||||
print("Done.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=(__doc__ or "").split("\n\n", 1)[0])
|
||||
parser.add_argument("--src", default=DEFAULT_REPO_ID, help="HF repo id or local Flux2 Klein snapshot directory")
|
||||
parser.add_argument("--dst", required=True, help="Output directory for the FastVideo Diffusers-style repo")
|
||||
parser.add_argument("--revision", default=None, help="Optional HF revision for --src repo id")
|
||||
parser.add_argument("--cache-dir", default=None, help="Optional Hugging Face cache directory")
|
||||
parser.add_argument("--max-shard-size", default="5GB", help="Maximum output shard size")
|
||||
parser.add_argument(
|
||||
"--transformer-source",
|
||||
choices=("auto", "raw", "diffusers"),
|
||||
default="auto",
|
||||
help="auto prefers flux-2-klein-4b.safetensors when present, else transformer/",
|
||||
)
|
||||
parser.add_argument("--overwrite", action="store_true", help="Replace a non-empty output directory")
|
||||
args = parser.parse_args()
|
||||
|
||||
convert(
|
||||
src=args.src,
|
||||
dst=args.dst,
|
||||
revision=args.revision,
|
||||
cache_dir=args.cache_dir,
|
||||
max_shard_size=args.max_shard_size,
|
||||
transformer_source=args.transformer_source,
|
||||
overwrite=args.overwrite,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,23 @@
|
||||
# Run with:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA fastvideo generate --config scripts/inference/inference_flux2_klein.yaml
|
||||
generator:
|
||||
model_path: black-forest-labs/FLUX.2-klein-4B
|
||||
engine:
|
||||
num_gpus: 1
|
||||
offload:
|
||||
dit: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
workload_type: t2i
|
||||
request:
|
||||
prompt: "a photo of a banana on a wooden table, studio lighting"
|
||||
sampling:
|
||||
seed: 0
|
||||
num_frames: 1
|
||||
height: 1024
|
||||
width: 1024
|
||||
num_inference_steps: 4
|
||||
guidance_scale: 1.0
|
||||
output:
|
||||
output_path: outputs/flux2-klein/
|
||||
@@ -23,6 +23,7 @@ tests/local_tests/<family>/
|
||||
|
||||
| Family | Workload | Dir |
|
||||
|---|---|---|
|
||||
| Flux2 | T2I | [`flux2/`](./flux2/) |
|
||||
| Hunyuan GameCraft | T2V / I2V | [`gamecraft/`](./gamecraft/) |
|
||||
| GEN3C | T2V | [`gen3c/`](./gen3c/) |
|
||||
| Kandinsky-5 | T2V | [`kandinsky5/`](./kandinsky5/) |
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
# Flux2 Local Tests
|
||||
|
||||
Local-only component and pipeline parity tests for the `flux2` FastVideo port.
|
||||
Compares FastVideo's Flux2 components and pipelines against the Diffusers
|
||||
reference. Flux2 Klein uses Qwen3 and the published
|
||||
`black-forest-labs/FLUX.2-klein-4B` checkpoint. Full Flux2 uses the Mistral3 /
|
||||
AutoProcessor text path and is activated with `FLUX2_FULL_MODEL_DIR`. Skipped in
|
||||
CI; CUDA required for activated parity runs.
|
||||
|
||||
## Reference Assets
|
||||
|
||||
| Field | Value |
|
||||
|---|---|
|
||||
| Model family | `flux2` |
|
||||
| Workload types | `T2I` |
|
||||
| Official reference | `diffusers.Flux2Pipeline`, `diffusers.Flux2KleinPipeline`, `diffusers.Flux2Transformer2DModel`, `diffusers.AutoencoderKLFlux2`, `transformers.Mistral3ForConditionalGeneration`, `transformers.Qwen3ForCausalLM` |
|
||||
| Local reference dir | `none` (Diffusers + transformers reference, no clone) |
|
||||
| Official commit/version | diffusers >= 0.38.0, transformers >= 4.52 |
|
||||
| HF weights | `black-forest-labs/FLUX.2-dev`, `black-forest-labs/FLUX.2-klein-4B`, `black-forest-labs/FLUX.2-klein-9B` |
|
||||
| HF revision | latest |
|
||||
| Local weights dir | Klein: `official_weights/black-forest-labs__FLUX.2-klein-4B` (env: `FLUX2_MODEL_DIR`); full: env `FLUX2_FULL_MODEL_DIR` |
|
||||
| Source layout | `diffusers` (native HF Diffusers format, no conversion needed) |
|
||||
| Needs conversion | No |
|
||||
|
||||
> Use only the env-var **name** for tokens (e.g., `HF_TOKEN`). Never paste a token value.
|
||||
|
||||
## Shared Environment Setup
|
||||
|
||||
Run from the FastVideo repo root in the same env used for FastVideo. The
|
||||
reference is the published Diffusers + transformers classes — no clone or
|
||||
upstream install is required beyond the FastVideo pins.
|
||||
|
||||
Do not change core dependency versions (`torch`, `transformers`, `flash-attn`,
|
||||
`triton`, CUDA packages) without explicit approval. The required Diffusers floor
|
||||
bump is recorded below.
|
||||
|
||||
## Official Environment Status
|
||||
|
||||
```text
|
||||
dependency_changes: diffusers>=0.38.0
|
||||
official_env_status: imports_ok
|
||||
private_dep_stubs: none
|
||||
blocked_on: none
|
||||
```
|
||||
|
||||
## Weight Setup
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"black-forest-labs/FLUX.2-klein-4B" \
|
||||
"official_weights/black-forest-labs__FLUX.2-klein-4B"
|
||||
```
|
||||
|
||||
## Tests in this directory
|
||||
|
||||
Port state is tracked in [`PORT_STATUS.md`](./PORT_STATUS.md).
|
||||
|
||||
```bash
|
||||
pytest tests/local_tests/flux2/ -v -s
|
||||
|
||||
FLUX2_MODEL_DIR=/path/to/black-forest-labs__FLUX.2-klein-4B \
|
||||
pytest tests/local_tests/pipelines/test_flux2_pipeline_smoke.py \
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py \
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py -v -s
|
||||
|
||||
FLUX2_FULL_MODEL_DIR=/path/to/black-forest-labs__FLUX.2-dev \
|
||||
pytest tests/local_tests/flux2/test_flux2_component_parity.py::test_flux2_mistral3_text_encoder_parity \
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py::test_flux2_full_transformer_guidance_parity \
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py::test_flux2_full_vae_encode_decode_parity \
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_smoke.py::test_flux2_full_pipeline_load_generate_smoke \
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py::test_flux2_full_pipeline_latent_parity -v -s
|
||||
```
|
||||
|
||||
| Component | Test | Concerns | Status |
|
||||
|---|---|---|---|
|
||||
| DiT transformer | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | Strict weight load + numerical output parity vs Diffusers, including 5D pipeline path | `PASSED on Modal L40S` |
|
||||
| Full DiT transformer | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | Strict full-weight load + embedded-guidance numerical output parity vs Diffusers | `PASSED on Modal L40S` |
|
||||
| VAE | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | Encode/decode exact parity vs Diffusers | `PASSED on Modal L40S` |
|
||||
| Full VAE | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | Full-weight encode/decode exact parity vs Diffusers | `PASSED on Modal L40S` |
|
||||
| Qwen3 text encoder | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | HF passthrough load + exact hidden-state parity | `PASSED on Modal L40S` |
|
||||
| Mistral3 text encoder | [`test_flux2_component_parity.py`](./test_flux2_component_parity.py) | Full Flux2 HF passthrough load + exact hidden-state parity | `PASSED on Modal L40S` |
|
||||
| Pipeline smoke | [`../pipelines/test_flux2_pipeline_smoke.py`](../pipelines/test_flux2_pipeline_smoke.py) | Import, registry, preset, config wiring; four-step latent generate | `PASSED on Modal L40S` |
|
||||
| Full pipeline smoke | [`../pipelines/test_flux2_pipeline_smoke.py`](../pipelines/test_flux2_pipeline_smoke.py) | Full Flux2 Mistral3/AutoProcessor wiring; short latent generate | `PASSED on Modal L40S:2` |
|
||||
| Pipeline parity | [`../pipelines/test_flux2_pipeline_parity.py`](../pipelines/test_flux2_pipeline_parity.py) | Four-step denoised latent parity vs `diffusers.Flux2KleinPipeline` | `PASSED on Modal L40S` |
|
||||
| Full pipeline parity | [`../pipelines/test_flux2_pipeline_parity.py`](../pipelines/test_flux2_pipeline_parity.py) | Short full Flux2 latent parity vs `diffusers.Flux2Pipeline` | `PASSED on Modal L40S:2, L40S:4, and H100:1` |
|
||||
| Pipeline TP2 parity | [`../pipelines/test_flux2_pipeline_parity.py`](../pipelines/test_flux2_pipeline_parity.py) | Two-worker tensor-parallel load/generate with `num_gpus=2`, `tp_size=2`, `sp_size=1` | `PASSED on Modal L40S:2` |
|
||||
| Pixel image comparison | Modal image runner | Same-prompt Diffusers vs FastVideo PNG generation and pixel metrics | `PASSED on Modal L40S` |
|
||||
|
||||
## Latest Remote Evidence
|
||||
|
||||
Modal L40S run `ap-5Zha6ev4NhKsIahsjdWiEb` applied patch
|
||||
`patches/flux2-local-cea4ed4d.patch` to commit
|
||||
`c23820e93d7b77d4113ca8fceac8ef3a19f572d3` with
|
||||
`FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`.
|
||||
|
||||
```text
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py -v -s: 3 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_smoke.py -v -s: 2 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py -v -s: 1 passed
|
||||
Flux2 modal final statuses: component=0 smoke=0 pipeline=0
|
||||
```
|
||||
|
||||
The parity tests print `assert_close` input means before every strict tensor
|
||||
comparison. The pipeline parity run reported zero max/mean/median diff for all
|
||||
four trajectory steps and final latents.
|
||||
|
||||
Additional Modal `L40S:2` allocation run `ap-sg5G52nqwBMDh1YN0IsoKd` applied
|
||||
patch `patches/flux2-local-l40s2-f708504d.patch` to the same commit and
|
||||
confirmed `torch.cuda.device_count() == 2`.
|
||||
|
||||
```text
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py -v -s: 3 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_smoke.py -v -s: 2 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py -v -s: 1 passed
|
||||
Flux2 modal L40S:2 statuses: component=0 smoke=0 pipeline=0
|
||||
```
|
||||
|
||||
The `L40S:2` run verifies the same parity suite in a two-GPU Modal allocation.
|
||||
These tests currently instantiate FastVideo with `num_gpus=1`, so this is not a
|
||||
tensor-parallel two-GPU parity test.
|
||||
|
||||
Additional Modal `L40S:4` allocation run `ap-tQEUnFr00uOvpMZdOngebQ` applied
|
||||
patch `patches/flux2-local-multil40s-cb92dbaa.patch` to the same commit and
|
||||
confirmed `torch.cuda.device_count() == 4`.
|
||||
|
||||
```text
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py -v -s: 3 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_smoke.py -v -s: 2 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py -v -s: 1 passed
|
||||
Flux2 modal L40S:4 statuses: component=0 smoke=0 pipeline=0
|
||||
```
|
||||
|
||||
Modal `L40S:2` tensor-parallel run `ap-szNgcJRiUv11lmmNFvjjPT` applied patch
|
||||
`patches/flux2-local-tp2-c850fcad.patch`, confirmed two visible L40S devices,
|
||||
and instantiated FastVideo with `num_gpus=2`, `tp_size=2`, `sp_size=1`,
|
||||
`executor_world_size=2`.
|
||||
|
||||
```text
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py::test_flux2_klein_pipeline_tensor_parallel_latent_parity -v -s: 1 passed
|
||||
TP2 final diff max=2.107544 mean=0.020110 median=0.015625
|
||||
TP2 abs-mean drift diffusers=1.319441 fastvideo=1.319235
|
||||
```
|
||||
|
||||
The TP2 test uses relaxed TP-specific bounds because BF16 tensor-parallel
|
||||
matmuls are not bit-exact with the single-GPU Diffusers reference. The strict
|
||||
single-GPU pipeline parity remains `atol=rtol=1e-4`.
|
||||
|
||||
Modal `L40S:1` image comparison run `ap-t0t6x3k15OoRhP56DYFcnj` generated a
|
||||
same-prompt Diffusers reference PNG and FastVideo PNG for prompt
|
||||
`a photo of a banana on a wooden table, studio lighting`, seed `0`,
|
||||
1024x1024, four steps, guidance `1.0`. Artifacts were downloaded locally under
|
||||
`outputs/flux2_image_compare/flux2_klein_seed0_files/`.
|
||||
|
||||
```text
|
||||
official_diffusers.png: 1024x1024 RGB
|
||||
fastvideo.png: 1024x1024 RGB
|
||||
pixel max_abs_diff=5 mean_abs_diff=0.480136 median_abs_diff=0.0 rmse=0.694312
|
||||
```
|
||||
|
||||
Modal `L40S:2` current-changes rerun `ap-CtccuhEHUwQy4Zv08nmrom` applied
|
||||
`/root/data/flux2_l40s2_current/runner/flux2-current.patch` to commit
|
||||
`69d22881a266306ad3bdbe820508ac17c13d2798` with
|
||||
`FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`, `FLUX2_TP_SIZE=2`, and two visible
|
||||
L40S devices.
|
||||
|
||||
```text
|
||||
tests/local_tests/flux2/test_flux2_component_parity.py -v -s: 3 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_smoke.py -v -s: 2 passed
|
||||
tests/local_tests/pipelines/test_flux2_pipeline_parity.py -v -s: 2 passed
|
||||
```
|
||||
|
||||
The pipeline parity file covered both strict single-GPU latent parity and true
|
||||
TP2 latent parity. The strict path reported zero max/mean/median diff for all
|
||||
four trajectory steps and final latents. The TP2 path instantiated FastVideo
|
||||
with `num_gpus=2`, `tp_size=2`, `sp_size=1`, and `executor_world_size=2`.
|
||||
|
||||
```text
|
||||
TP2 final diff max=2.107544 mean=0.020110 median=0.015625
|
||||
TP2 abs-mean drift diffusers=1.319441 fastvideo=1.319235
|
||||
```
|
||||
|
||||
The same Modal run also regenerated the same-prompt Diffusers and FastVideo PNGs
|
||||
with the current changes. Artifacts were downloaded locally under
|
||||
`flux2_image_compare_results_l40s2_current/`.
|
||||
|
||||
```text
|
||||
official_diffusers.png: 1024x1024 RGB
|
||||
fastvideo.png: 1024x1024 RGB
|
||||
comparison_grid.png: 3072x1058 RGB
|
||||
pixel max_abs_diff=5 mean_abs_diff=0.480136 median_abs_diff=0.0 rmse=0.694312
|
||||
```
|
||||
|
||||
Full Flux2 weight probe:
|
||||
|
||||
- Modal run `ap-70hWVB2eZAE1JWw0Im4E2T` checked
|
||||
`/root/data/official_weights` and found only
|
||||
`/root/data/official_weights/black-forest-labs__FLUX.2-klein-4B`.
|
||||
- Modal run `ap-pp2v3Iu3NvJrkiMVdsIk0x` confirmed the same with
|
||||
`ls -la /root/data/official_weights`.
|
||||
- Modal run `ap-erDuzyLLnsd3OXKaVAhlUW` rechecked the current volume and still
|
||||
found only `/root/data/official_weights/black-forest-labs__FLUX.2-klein-4B`.
|
||||
- Local env-var name probe found `HF_TOKEN`, `HUGGINGFACE_HUB_TOKEN`, and
|
||||
`HF_API_KEY` absent, so the launcher cannot pass a gated HF token through to
|
||||
Modal.
|
||||
|
||||
The full `FLUX.2-dev` checkpoint was later staged and committed to the Modal
|
||||
`hf-model-weights` volume by app `ap-0SFRuDEG24nHlGKWhZwCcj`; the staged path is
|
||||
`/root/data/official_weights/black-forest-labs__FLUX.2-dev`.
|
||||
|
||||
Current-session Modal evidence:
|
||||
|
||||
- Modal run `ap-tIeZbOsqsjMfN5qXJoz0WG` applied the current local patch after
|
||||
the launcher venv fix and passed:
|
||||
`test_flux2_full_typed_surface_preflight` and
|
||||
`test_flux2_full_text_stage_uses_mistral3_format_and_embedded_guidance`.
|
||||
- Modal run `ap-oetobVLKHRX7uBkZ1ZTo7X` applied the current local patch, ran
|
||||
`examples/inference/basic/basic_flux2_klein.py`, and verified
|
||||
`outputs/flux2/flux2_klein_example.png` as `1024x1024 RGB`.
|
||||
- Modal run `ap-d3yQLJ2jwAYewcCr5pyBYJ` passed full Mistral3 hidden-state
|
||||
parity with exact zero diff.
|
||||
- Modal app logs for `ap-i6ZLmxsptBW49OP1DC38ZE` confirmed full transformer
|
||||
CPU BF16 parity passed with exact zero diff.
|
||||
- Modal run `ap-5Nm5rPHf0Dph5M6En9azVT` passed full VAE encode/decode parity.
|
||||
- Modal `L40S:2` run `ap-XPH4aM4LZxKhHJz7IDSMRg` passed full pipeline smoke at
|
||||
`128x128`, one step, `max_sequence_length=64`, `tp_size=2`, `sp_size=1`.
|
||||
- Modal `L40S:2` run `ap-nVr8tDtQ0lVwviZDm6rIjH` passed full pipeline latent
|
||||
parity. Final latent diff max `0.062500`, mean `0.008581`, median
|
||||
`0.007812`.
|
||||
- Modal `L40S:2` diagnostic run `ap-nhB228LmgVjGP6oRuCYSPT` reproduced the
|
||||
full pipeline latent parity diff and printed quantization details: max
|
||||
`0.062500`, mean `0.008581`, median `0.007812`; only `9/8192` latent entries
|
||||
hit the max bucket.
|
||||
- Modal `L40S:4` run `ap-KKOzv9THDmTYt3gevhQkVR` passed full pipeline latent
|
||||
parity with true TP4 (`num_gpus=4`, `tp_size=4`, `sp_size=1`). Final latent
|
||||
diff max stayed `0.062500`, mean was `0.007946`, and median was `0.007812`.
|
||||
- Modal `L40S:4` diagnostic run `ap-SlhykTTLxzsYxnsCUbCH7S` reproduced the TP4
|
||||
result with max `0.062500`, mean `0.007946`, median `0.007812`; only `6/8192`
|
||||
latent entries hit the max bucket.
|
||||
- Modal `L40S:2` input-variant diagnostic run `ap-5N1yHy8udeZKslrQJTFHYX` used
|
||||
`FLUX2_FULL_RUN_INPUT_VARIANTS=1` to change prompt/seed. Changed prompt with
|
||||
seed `0` produced max `0.062500`, mean `0.006670`, median `0.003906`, and
|
||||
`5/8192` max-bucket entries. Default prompt with seed `123` produced max
|
||||
`0.062500`, mean `0.008709`, median `0.007812`, and `22/8192` max-bucket
|
||||
entries.
|
||||
- Modal `L40S:4` input-variant diagnostic run `ap-TvJQ5cvNculTxJbTJAPNzx`
|
||||
passed the same cases. Changed prompt with seed `0` produced max `0.062500`,
|
||||
mean `0.006709`, median `0.003906`, and `3/8192` max-bucket entries. Default
|
||||
prompt with seed `123` produced max `0.062500`, mean `0.008820`, median
|
||||
`0.007812`, and `8/8192` max-bucket entries.
|
||||
- Modal `L40S:1` run `ap-A8mRqCmPjUEs3kZJ3j8twH` attempted the same current
|
||||
full pipeline latent parity command with TP1. Diffusers produced reference
|
||||
latents, but FastVideo OOMed while loading the full transformer before parity
|
||||
comparison.
|
||||
- Modal `L40S:1` setup probe `ap-FfM5ggOjtmc1oYK93T9Bvd` passed
|
||||
`test_flux2_full_pipeline_setup_matches_diffusers` with exact zero diff for
|
||||
Mistral3 prompt embeddings, text ids, raw/packed latents, image ids,
|
||||
timesteps, guidance scaling, and packed-vs-5D scheduler stepping.
|
||||
- Modal `L40S:1` direct CUDA full-transformer component attempt
|
||||
`ap-LJm3FpKIJJzf1mosv7qfmd` OOMed before forward while moving the Diffusers
|
||||
transformer to GPU, so TP1/single-GPU CUDA full-transformer evidence remains
|
||||
infeasible on one L40S.
|
||||
- Modal `H100:1` run `ap-HOpne8l9NbpWwU9qhUXYxT` used the updated
|
||||
`fastvideo/tests/modal/launch_l40s_job.py --gpu-type H100` path and passed
|
||||
`test_flux2_full_pipeline_latent_parity` with `FLUX2_FULL_NUM_GPUS=1`,
|
||||
`FLUX2_FULL_TP_SIZE=1`, and `FLUX2_FULL_SP_SIZE=1`. The trajectory step 0 and
|
||||
final packed latent diffs were exactly zero: max `0.000000`, mean `0.000000`,
|
||||
median `0.000000`.
|
||||
- Modal `H100:1` run `ap-Gy6MRyZduxTHihgxiWWL13` generated full Flux2
|
||||
Diffusers/FastVideo image comparison artifacts at `1024x1024`, four steps,
|
||||
guidance `4.0`, max sequence length `64`. Local artifacts are in
|
||||
`flux2_full_image_compare_20260526_h100_full_t2i_files/`; pixel metrics:
|
||||
max absolute diff `14`, mean absolute diff `0.658212`, median `1.0`, RMSE
|
||||
`0.848854`.
|
||||
- Modal `L40S:2` run `ap-xJ1MnjIQz79KPD4X9EOTLY` ran
|
||||
`examples/inference/basic/basic_flux2.py` and verified the generated PNG as
|
||||
`(128, 128) RGB`.
|
||||
|
||||
## Scope Notes
|
||||
|
||||
- **Validated**: Flux2 Klein (distilled, 4-step, Qwen3 text encoder, no guidance).
|
||||
End-to-end latent inference matches Diffusers exactly for the four-step Klein
|
||||
pipeline parity prompt.
|
||||
- **Validated**: Full Flux2 T2I (`Flux2Pipeline`) uses
|
||||
Mistral3/AutoProcessor text conditioning and treats `guidance_scale` as
|
||||
embedded transformer guidance. Full component parity, pipeline smoke, latent
|
||||
parity, and example generation passed on Modal with `FLUX2_FULL_MODEL_DIR`.
|
||||
- **Investigated**: The full TP latent max diff `0.062500` is not caused by
|
||||
prompt setup, latent packing, ids, timesteps, guidance scaling, or scheduler
|
||||
layout. Dedicated setup parity is bit-exact, and prompt/seed changes preserve
|
||||
the same worst-case bucket; the remaining diff is rare BF16-grid TP denoiser
|
||||
drift.
|
||||
- **Validated**: Full Flux2 single-GPU pipeline parity is exact on Modal H100:1.
|
||||
The nonzero full pipeline diff is therefore specific to tensor-parallel
|
||||
execution, not the full model wiring.
|
||||
- **Deferred**: Full Flux2 image conditioning and caption upsampling are not
|
||||
claimed by this port.
|
||||
|
||||
## Review Notes
|
||||
|
||||
- Required before handoff: non-skip PASS for each component parity test,
|
||||
including reused components that own weights or numerical behavior.
|
||||
- Pipeline parity may start as pending coverage; final handoff requires non-skip
|
||||
PASS or an explicit blocker accepted via the escape-hatch process.
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,471 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Trace Klein Flux2 dense-vs-TP2 transformer drift.
|
||||
|
||||
Run only on Modal/GPU, for example:
|
||||
|
||||
FLUX2_MODEL_DIR=/root/data/official_weights/black-forest-labs__FLUX.2-klein-4B \
|
||||
python -m torch.distributed.run --nproc_per_node=2 \
|
||||
tests/local_tests/flux2/debug_flux2_klein_tp_trace.py
|
||||
|
||||
Rank 0 runs a dense Diffusers reference forward, then both ranks run the
|
||||
FastVideo TP2 forward with identical tensors. Rank 0 compares captured top-level
|
||||
module outputs and writes diff-friendly summaries to /tmp/opencode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
TRACE_DIR = Path("/tmp/opencode")
|
||||
REF_LOG = TRACE_DIR / "flux2_klein_tp2_ref_layers.log"
|
||||
FV_LOG = TRACE_DIR / "flux2_klein_tp2_fv_layers.log"
|
||||
DIFF_LOG = TRACE_DIR / "flux2_klein_tp2_diff.log"
|
||||
|
||||
|
||||
def _load_json(path: Path) -> dict[str, Any]:
|
||||
with path.open("r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _collect_safetensors_paths(model_dir: Path) -> list[str]:
|
||||
paths = sorted(str(path) for path in model_dir.glob("*.safetensors"))
|
||||
if not paths:
|
||||
raise FileNotFoundError(f"No safetensors files found under {model_dir}")
|
||||
return paths
|
||||
|
||||
|
||||
def _collect_safetensors_keys(paths: Iterable[str]) -> set[str]:
|
||||
from safetensors.torch import safe_open
|
||||
|
||||
keys: set[str] = set()
|
||||
for path in paths:
|
||||
with safe_open(path, framework="pt", device="cpu") as f:
|
||||
keys.update(f.keys())
|
||||
return keys
|
||||
|
||||
|
||||
def _load_tensor_from_safetensors(paths: Iterable[str], key: str) -> torch.Tensor:
|
||||
from safetensors.torch import safe_open
|
||||
|
||||
for path in paths:
|
||||
with safe_open(path, framework="pt", device="cpu") as f:
|
||||
if key in f.keys():
|
||||
return f.get_tensor(key)
|
||||
raise KeyError(f"Could not find {key!r} in safetensors files")
|
||||
|
||||
|
||||
class _DenseLinearTuple(torch.nn.Module):
|
||||
def __init__(self, weight: torch.Tensor, bias: torch.Tensor | None = None):
|
||||
super().__init__()
|
||||
self.weight = torch.nn.Parameter(weight, requires_grad=False)
|
||||
self.bias = None if bias is None else torch.nn.Parameter(bias, requires_grad=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, None]:
|
||||
return torch.nn.functional.linear(x, self.weight, self.bias), None
|
||||
|
||||
|
||||
def _record_tensor(trace: dict[str, torch.Tensor], name: str, tensor: torch.Tensor) -> None:
|
||||
trace[name] = tensor.detach().float().cpu()
|
||||
|
||||
|
||||
def _record_output(trace: dict[str, torch.Tensor], name: str, output: Any) -> None:
|
||||
if torch.is_tensor(output):
|
||||
_record_tensor(trace, name, output)
|
||||
return
|
||||
if isinstance(output, tuple) and len(output) == 2 and torch.is_tensor(output[0]) and output[1] is None:
|
||||
_record_tensor(trace, name, output[0])
|
||||
return
|
||||
if isinstance(output, (tuple, list)):
|
||||
for index, item in enumerate(output):
|
||||
_record_output(trace, f"{name}[{index}]", item)
|
||||
|
||||
|
||||
def _get_submodule(model: torch.nn.Module, name: str) -> torch.nn.Module | None:
|
||||
current: torch.nn.Module = model
|
||||
for part in name.split("."):
|
||||
if part.isdigit() and isinstance(current, torch.nn.ModuleList):
|
||||
current = current[int(part)]
|
||||
continue
|
||||
child = getattr(current, part, None)
|
||||
if not isinstance(child, torch.nn.Module):
|
||||
return None
|
||||
current = child
|
||||
return current
|
||||
|
||||
|
||||
def _set_submodule(model: torch.nn.Module, name: str, module: torch.nn.Module) -> None:
|
||||
parts = name.split(".")
|
||||
parent_name = ".".join(parts[:-1])
|
||||
child_name = parts[-1]
|
||||
parent = _get_submodule(model, parent_name) if parent_name else model
|
||||
if parent is None:
|
||||
raise AttributeError(f"Could not find parent module for {name!r}")
|
||||
if child_name.isdigit() and isinstance(parent, torch.nn.ModuleList):
|
||||
parent[int(child_name)] = module
|
||||
else:
|
||||
setattr(parent, child_name, module)
|
||||
|
||||
|
||||
def _patch_dense_linear_tuple(
|
||||
model: torch.nn.Module,
|
||||
weight_paths: Iterable[str],
|
||||
available_keys: set[str],
|
||||
module_name: str,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> None:
|
||||
weight_key = f"{module_name}.weight"
|
||||
bias_key = f"{module_name}.bias"
|
||||
weight = _load_tensor_from_safetensors(weight_paths, weight_key).to(device=device, dtype=dtype)
|
||||
bias = None
|
||||
if bias_key in available_keys:
|
||||
bias = _load_tensor_from_safetensors(weight_paths, bias_key).to(device=device, dtype=dtype)
|
||||
_set_submodule(model, module_name, _DenseLinearTuple(weight, bias).to(device=device))
|
||||
print(f"[debug] Patched {module_name} to replicated dense F.linear")
|
||||
|
||||
|
||||
def _attach_hooks(
|
||||
model: torch.nn.Module,
|
||||
names: Iterable[str],
|
||||
trace: dict[str, torch.Tensor],
|
||||
) -> list[torch.utils.hooks.RemovableHandle]:
|
||||
handles: list[torch.utils.hooks.RemovableHandle] = []
|
||||
for name in names:
|
||||
module = _get_submodule(model, name)
|
||||
if module is None:
|
||||
print(f"[trace] missing module {name}")
|
||||
continue
|
||||
|
||||
def _hook(_module, _inputs, output, *, hook_name=name):
|
||||
_record_output(trace, hook_name, output)
|
||||
|
||||
handles.append(module.register_forward_hook(_hook))
|
||||
return handles
|
||||
|
||||
|
||||
def _trace_names(num_double: int, num_single: int) -> list[str]:
|
||||
names = [
|
||||
"time_guidance_embed",
|
||||
"double_stream_modulation_img",
|
||||
"double_stream_modulation_txt",
|
||||
"single_stream_modulation",
|
||||
"x_embedder",
|
||||
"context_embedder",
|
||||
]
|
||||
names.extend(f"transformer_blocks.{idx}" for idx in range(num_double))
|
||||
names.extend(f"single_transformer_blocks.{idx}" for idx in range(num_single))
|
||||
names.extend(["norm_out", "proj_out"])
|
||||
drill_double = os.getenv("FLUX2_KLEIN_TP_TRACE_DRILL_DOUBLE_BLOCK", "")
|
||||
if drill_double:
|
||||
base = f"transformer_blocks.{int(drill_double)}"
|
||||
names.extend(
|
||||
f"{base}.{suffix}"
|
||||
for suffix in (
|
||||
"norm1",
|
||||
"norm1_context",
|
||||
"attn.to_q",
|
||||
"attn.to_k",
|
||||
"attn.to_v",
|
||||
"attn.add_q_proj",
|
||||
"attn.add_k_proj",
|
||||
"attn.add_v_proj",
|
||||
"attn.norm_q",
|
||||
"attn.norm_k",
|
||||
"attn.norm_added_q",
|
||||
"attn.norm_added_k",
|
||||
"attn.to_add_out",
|
||||
"attn.to_out.0",
|
||||
"attn",
|
||||
"norm2",
|
||||
"ff.linear_in",
|
||||
"ff.act_fn",
|
||||
"ff.linear_out",
|
||||
"ff",
|
||||
"norm2_context",
|
||||
"ff_context.linear_in",
|
||||
"ff_context.act_fn",
|
||||
"ff_context.linear_out",
|
||||
"ff_context",
|
||||
)
|
||||
)
|
||||
return names
|
||||
|
||||
|
||||
def _write_trace(path: Path, trace: dict[str, torch.Tensor]) -> None:
|
||||
with path.open("w", encoding="utf-8") as f:
|
||||
for name, tensor in trace.items():
|
||||
t = tensor.float()
|
||||
f.write(
|
||||
f"{name} {tuple(t.shape)} "
|
||||
f"{t.abs().mean().item():.8f} {t.sum().item():.8f} "
|
||||
f"{t.min().item():.8f} {t.max().item():.8f}\n"
|
||||
)
|
||||
|
||||
|
||||
def _compare_traces(
|
||||
ref_trace: dict[str, torch.Tensor],
|
||||
fv_trace: dict[str, torch.Tensor],
|
||||
) -> None:
|
||||
first = None
|
||||
with DIFF_LOG.open("w", encoding="utf-8") as f:
|
||||
for name, ref_tensor in ref_trace.items():
|
||||
fv_tensor = fv_trace.get(name)
|
||||
if fv_tensor is None:
|
||||
f.write(f"{name} missing_on_fastvideo\n")
|
||||
if first is None:
|
||||
first = (name, "missing")
|
||||
continue
|
||||
if tuple(ref_tensor.shape) != tuple(fv_tensor.shape):
|
||||
f.write(f"{name} shape ref={tuple(ref_tensor.shape)} fv={tuple(fv_tensor.shape)}\n")
|
||||
if first is None:
|
||||
first = (name, "shape")
|
||||
continue
|
||||
diff = (ref_tensor - fv_tensor).abs()
|
||||
max_diff = diff.max().item()
|
||||
mean_diff = diff.mean().item()
|
||||
median_diff = diff.median().item()
|
||||
f.write(
|
||||
f"{name} max={max_diff:.8f} mean={mean_diff:.8f} "
|
||||
f"median={median_diff:.8f} ref_abs={ref_tensor.abs().mean().item():.8f} "
|
||||
f"fv_abs={fv_tensor.abs().mean().item():.8f}\n"
|
||||
)
|
||||
if first is None and max_diff > 0:
|
||||
first = (name, f"max={max_diff:.8f} mean={mean_diff:.8f}")
|
||||
if first is None:
|
||||
print("[trace] no divergence across captured tensors")
|
||||
else:
|
||||
print(f"[trace] first divergence: {first[0]} {first[1]}")
|
||||
print(f"[trace] wrote {REF_LOG}")
|
||||
print(f"[trace] wrote {FV_LOG}")
|
||||
print(f"[trace] wrote {DIFF_LOG}")
|
||||
if DIFF_LOG.exists():
|
||||
print(DIFF_LOG.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def _build_inputs(
|
||||
*,
|
||||
batch: int,
|
||||
img_h: int,
|
||||
img_w: int,
|
||||
txt_len: int,
|
||||
in_channels: int,
|
||||
joint_dim: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch.manual_seed(0)
|
||||
seq_len = img_h * img_w
|
||||
hidden = torch.randn(batch, seq_len, in_channels, dtype=torch.float32)
|
||||
encoder = torch.randn(batch, txt_len, joint_dim, dtype=torch.float32)
|
||||
timestep = torch.tensor([0.5], dtype=torch.float32)
|
||||
guidance = torch.zeros_like(timestep)
|
||||
txt_ids = torch.cartesian_prod(
|
||||
torch.arange(1),
|
||||
torch.arange(1),
|
||||
torch.arange(1),
|
||||
torch.arange(txt_len),
|
||||
)
|
||||
img_ids = torch.cartesian_prod(
|
||||
torch.arange(1),
|
||||
torch.arange(img_h),
|
||||
torch.arange(img_w),
|
||||
torch.arange(1),
|
||||
)
|
||||
return hidden, encoder, timestep, guidance, txt_ids, img_ids
|
||||
|
||||
|
||||
def _run_reference(
|
||||
transformer_dir: Path,
|
||||
cfg: dict[str, Any],
|
||||
names: list[str],
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
inputs: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
) -> dict[str, torch.Tensor]:
|
||||
from diffusers import Flux2Transformer2DModel as RefTransformer
|
||||
|
||||
trace: dict[str, torch.Tensor] = {}
|
||||
ref = RefTransformer.from_pretrained(
|
||||
str(transformer_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=False,
|
||||
).eval().to(device)
|
||||
|
||||
class _ZeroGuidance(torch.nn.Module):
|
||||
def __init__(self, embedding_dim: int):
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
def forward(self, guidance_proj: torch.Tensor) -> torch.Tensor:
|
||||
return torch.zeros(
|
||||
guidance_proj.shape[0],
|
||||
self.embedding_dim,
|
||||
device=guidance_proj.device,
|
||||
dtype=guidance_proj.dtype,
|
||||
)
|
||||
|
||||
ref.time_guidance_embed.guidance_embedder = _ZeroGuidance(int(cfg["num_attention_heads"]) * int(cfg["attention_head_dim"]))
|
||||
handles = _attach_hooks(ref, names, trace)
|
||||
hidden, encoder, timestep, guidance, txt_ids, img_ids = inputs
|
||||
try:
|
||||
with torch.no_grad():
|
||||
output = ref(
|
||||
hidden_states=hidden.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=encoder.to(device=device, dtype=dtype),
|
||||
timestep=timestep.to(device=device, dtype=dtype),
|
||||
img_ids=img_ids.to(device=device),
|
||||
txt_ids=txt_ids.to(device=device),
|
||||
guidance=guidance.to(device=device, dtype=dtype),
|
||||
return_dict=False,
|
||||
)[0]
|
||||
_record_tensor(trace, "output", output)
|
||||
finally:
|
||||
for handle in handles:
|
||||
handle.remove()
|
||||
del ref
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
_ = cfg
|
||||
return trace
|
||||
|
||||
|
||||
def _run_fastvideo_tp(
|
||||
transformer_dir: Path,
|
||||
cfg: dict[str, Any],
|
||||
names: list[str],
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
inputs: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
rank: int,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("Flux2Transformer2DModel")
|
||||
dit_cfg = Flux2Config()
|
||||
dit_cfg.update_model_arch(dict(cfg))
|
||||
update_fn = getattr(dit_cfg.arch_config, "update_from_weight_keys", None)
|
||||
weight_paths = _collect_safetensors_paths(transformer_dir)
|
||||
weight_keys = _collect_safetensors_keys(weight_paths)
|
||||
if callable(update_fn):
|
||||
update_fn(weight_keys)
|
||||
|
||||
fv = maybe_load_fsdp_model(
|
||||
model_cls=fv_cls,
|
||||
init_params={"config": dit_cfg, "hf_config": dict(cfg)},
|
||||
weight_dir_list=weight_paths,
|
||||
device=device,
|
||||
hsdp_replicate_dim=1,
|
||||
hsdp_shard_dim=1,
|
||||
strict=True,
|
||||
cpu_offload=False,
|
||||
fsdp_inference=False,
|
||||
default_dtype=dtype,
|
||||
param_dtype=dtype,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
training_mode=False,
|
||||
pin_cpu_memory=False,
|
||||
).eval()
|
||||
|
||||
dense_modules: list[str] = []
|
||||
if os.getenv("FLUX2_KLEIN_TP_TRACE_PATCH_CONTEXT_DENSE", "0") == "1":
|
||||
dense_modules.append("context_embedder")
|
||||
if os.getenv("FLUX2_KLEIN_TP_TRACE_PATCH_BLOCK0_FF_CONTEXT_OUT_DENSE", "0") == "1":
|
||||
dense_modules.append("transformer_blocks.0.ff_context.linear_out")
|
||||
dense_modules.extend(
|
||||
name.strip()
|
||||
for name in os.getenv("FLUX2_KLEIN_TP_TRACE_DENSE_LINEAR_MODULES", "").split(",")
|
||||
if name.strip()
|
||||
)
|
||||
for module_name in dict.fromkeys(dense_modules):
|
||||
_patch_dense_linear_tuple(fv, weight_paths, weight_keys, module_name, device, dtype)
|
||||
|
||||
trace: dict[str, torch.Tensor] = {}
|
||||
handles = _attach_hooks(fv, names, trace) if rank == 0 else []
|
||||
hidden, encoder, timestep, guidance, txt_ids, img_ids = inputs
|
||||
try:
|
||||
with torch.no_grad():
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
output = fv(
|
||||
hidden_states=hidden.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=encoder.to(device=device, dtype=dtype),
|
||||
timestep=timestep.to(device=device, dtype=dtype),
|
||||
img_ids=img_ids.to(device=device),
|
||||
txt_ids=txt_ids.to(device=device),
|
||||
guidance=guidance.to(device=device, dtype=dtype),
|
||||
)
|
||||
if rank == 0:
|
||||
_record_tensor(trace, "output", output)
|
||||
finally:
|
||||
for handle in handles:
|
||||
handle.remove()
|
||||
del fv
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
return trace
|
||||
|
||||
|
||||
def main() -> None:
|
||||
model_dir = Path(os.getenv("FLUX2_MODEL_DIR", ""))
|
||||
if not model_dir.exists():
|
||||
raise RuntimeError("Set FLUX2_MODEL_DIR to the Klein checkpoint directory")
|
||||
transformer_dir = model_dir / "transformer"
|
||||
cfg = _load_json(transformer_dir / "config.json")
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
rank = int(os.environ.get("RANK", "0"))
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
dtype = torch.bfloat16
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
|
||||
from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=2, sp_size=1)
|
||||
|
||||
num_double = int(cfg.get("num_layers", 19))
|
||||
num_single = int(cfg.get("num_single_layers", 38))
|
||||
names = _trace_names(num_double, num_single)
|
||||
inputs = _build_inputs(
|
||||
batch=1,
|
||||
img_h=8,
|
||||
img_w=8,
|
||||
txt_len=16,
|
||||
in_channels=int(cfg["in_channels"]),
|
||||
joint_dim=int(cfg["joint_attention_dim"]),
|
||||
)
|
||||
|
||||
TRACE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
ref_trace: dict[str, torch.Tensor] = {}
|
||||
if rank == 0:
|
||||
ref_trace = _run_reference(transformer_dir, cfg, names, dtype, device, inputs)
|
||||
_write_trace(REF_LOG, ref_trace)
|
||||
|
||||
dist.barrier()
|
||||
fv_trace = _run_fastvideo_tp(transformer_dir, cfg, names, dtype, device, inputs, rank)
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
_write_trace(FV_LOG, fv_trace)
|
||||
_compare_traces(ref_trace, fv_trace)
|
||||
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,350 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Generate full Flux2 Diffusers/FastVideo image comparison artifacts.
|
||||
|
||||
This is intended for Modal GPU runs. The parent process tries requested square
|
||||
resolutions in order; each attempt runs in a fresh child process so a CUDA OOM
|
||||
can fall back cleanly to the next resolution.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
|
||||
DEFAULT_PROMPT = "a photo of a banana on a wooden table, studio lighting"
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Generate full Flux2 image comparison artifacts.")
|
||||
parser.add_argument("--model-dir", default=os.getenv("FLUX2_FULL_MODEL_DIR", ""))
|
||||
parser.add_argument("--output-root", default="/root/data/flux2_full_image_compare")
|
||||
parser.add_argument("--run-name", default="20260526_h100_full_t2i")
|
||||
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--sizes", default="1024,768,512,256")
|
||||
parser.add_argument("--steps", type=int, default=4)
|
||||
parser.add_argument("--guidance-scale", type=float, default=4.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=64)
|
||||
parser.add_argument("--child-size", type=int, default=0)
|
||||
parser.add_argument("--child-mode", choices=("diffusers", "fastvideo"), default="")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _load_prompt_embeds(
|
||||
model_dir: Path,
|
||||
prompt: str,
|
||||
dtype: torch.dtype,
|
||||
max_sequence_length: int,
|
||||
) -> torch.Tensor:
|
||||
from diffusers import Flux2Pipeline
|
||||
|
||||
prompt_pipe = Flux2Pipeline.from_pretrained(
|
||||
str(model_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
try:
|
||||
with torch.no_grad():
|
||||
return prompt_pipe._get_mistral_3_small_prompt_embeds( # noqa: SLF001
|
||||
text_encoder=prompt_pipe.text_encoder,
|
||||
tokenizer=prompt_pipe.tokenizer,
|
||||
prompt=[prompt],
|
||||
device=torch.device("cpu"),
|
||||
max_sequence_length=max_sequence_length,
|
||||
system_message=prompt_pipe.system_message,
|
||||
hidden_states_layers=(10, 20, 30),
|
||||
)
|
||||
finally:
|
||||
del prompt_pipe
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def _generate_diffusers_image(
|
||||
model_dir: Path,
|
||||
output_path: Path,
|
||||
*,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
size: int,
|
||||
steps: int,
|
||||
guidance_scale: float,
|
||||
max_sequence_length: int,
|
||||
) -> Image.Image:
|
||||
from diffusers import Flux2Pipeline
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
|
||||
prompt_embeds = _load_prompt_embeds(model_dir, prompt, dtype, max_sequence_length)
|
||||
max_memory = {idx: "78GiB" for idx in range(torch.cuda.device_count())}
|
||||
pipe = Flux2Pipeline.from_pretrained(
|
||||
str(model_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
text_encoder=None,
|
||||
tokenizer=None,
|
||||
device_map="balanced",
|
||||
max_memory=max_memory,
|
||||
)
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
try:
|
||||
with torch.no_grad():
|
||||
output = pipe(
|
||||
prompt=None,
|
||||
prompt_embeds=prompt_embeds.to(device=device, dtype=dtype),
|
||||
height=size,
|
||||
width=size,
|
||||
num_inference_steps=steps,
|
||||
guidance_scale=guidance_scale,
|
||||
max_sequence_length=max_sequence_length,
|
||||
generator=torch.Generator(device="cpu").manual_seed(seed),
|
||||
output_type="pil",
|
||||
return_dict=True,
|
||||
)
|
||||
image = output.images[0].convert("RGB")
|
||||
image.save(output_path)
|
||||
return image
|
||||
finally:
|
||||
del pipe
|
||||
del prompt_embeds
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def _generate_fastvideo_image(
|
||||
model_dir: Path,
|
||||
output_path: Path,
|
||||
*,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
size: int,
|
||||
steps: int,
|
||||
guidance_scale: float,
|
||||
max_sequence_length: int,
|
||||
) -> Image.Image:
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(model_dir),
|
||||
num_gpus=1,
|
||||
tp_size=1,
|
||||
sp_size=1,
|
||||
workload_type="t2i",
|
||||
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="Flux2Pipeline",
|
||||
)
|
||||
try:
|
||||
sampling = SamplingParam.from_pretrained(str(model_dir))
|
||||
sampling.prompt = prompt
|
||||
sampling.height = size
|
||||
sampling.width = size
|
||||
sampling.num_frames = 1
|
||||
sampling.fps = 1
|
||||
sampling.num_inference_steps = steps
|
||||
sampling.guidance_scale = guidance_scale
|
||||
sampling.max_sequence_length = max_sequence_length
|
||||
sampling.seed = seed
|
||||
sampling.output_path = str(output_path)
|
||||
sampling.save_video = True
|
||||
sampling.return_frames = False
|
||||
generator.generate_video(prompt, sampling_param=sampling, output_path=str(output_path))
|
||||
finally:
|
||||
generator.shutdown()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return Image.open(output_path).convert("RGB")
|
||||
|
||||
|
||||
def _save_comparison(
|
||||
output_dir: Path,
|
||||
upstream: Image.Image,
|
||||
fastvideo: Image.Image,
|
||||
metadata: dict[str, Any],
|
||||
) -> None:
|
||||
upstream_arr = np.asarray(upstream.convert("RGB"), dtype=np.int16)
|
||||
fastvideo_arr = np.asarray(fastvideo.convert("RGB"), dtype=np.int16)
|
||||
diff = np.abs(upstream_arr - fastvideo_arr).astype(np.uint8)
|
||||
max_diff = int(diff.max())
|
||||
metrics = {
|
||||
**metadata,
|
||||
"max_abs_diff": max_diff,
|
||||
"mean_abs_diff": float(diff.mean()),
|
||||
"median_abs_diff": float(np.median(diff)),
|
||||
"rmse": float(np.sqrt(np.mean((upstream_arr - fastvideo_arr).astype(np.float32) ** 2))),
|
||||
}
|
||||
|
||||
Image.fromarray(diff, mode="RGB").save(output_dir / "abs_diff.png")
|
||||
scale = 1 if max_diff == 0 else min(255.0 / max_diff, 64.0)
|
||||
Image.fromarray(np.clip(diff.astype(np.float32) * scale, 0, 255).astype(np.uint8), mode="RGB").save(
|
||||
output_dir / "abs_diff_scaled.png"
|
||||
)
|
||||
|
||||
label_height = 34
|
||||
gap = 8
|
||||
side_by_side = Image.new(
|
||||
"RGB",
|
||||
(upstream.width + fastvideo.width + gap, upstream.height + label_height),
|
||||
"white",
|
||||
)
|
||||
draw = ImageDraw.Draw(side_by_side)
|
||||
draw.text((0, 8), "Diffusers", fill=(0, 0, 0))
|
||||
draw.text((upstream.width + gap, 8), "FastVideo", fill=(0, 0, 0))
|
||||
side_by_side.paste(upstream, (0, label_height))
|
||||
side_by_side.paste(fastvideo, (upstream.width + gap, label_height))
|
||||
side_by_side.save(output_dir / "side_by_side.png")
|
||||
|
||||
with (output_dir / "metrics.json").open("w", encoding="utf-8") as f:
|
||||
json.dump(metrics, f, indent=2, sort_keys=True)
|
||||
with (output_dir / "README.txt").open("w", encoding="utf-8") as f:
|
||||
f.write(
|
||||
"Full Flux2 H100 image comparison\n"
|
||||
f"prompt: {metadata['prompt']}\n"
|
||||
f"seed: {metadata['seed']}\n"
|
||||
f"size: {metadata['size']}x{metadata['size']}\n"
|
||||
f"steps: {metadata['steps']}\n"
|
||||
f"guidance_scale: {metadata['guidance_scale']}\n"
|
||||
f"max_sequence_length: {metadata['max_sequence_length']}\n"
|
||||
f"max_abs_diff: {metrics['max_abs_diff']}\n"
|
||||
f"mean_abs_diff: {metrics['mean_abs_diff']}\n"
|
||||
f"median_abs_diff: {metrics['median_abs_diff']}\n"
|
||||
f"rmse: {metrics['rmse']}\n"
|
||||
)
|
||||
|
||||
|
||||
def _attempt_dir(args: argparse.Namespace, size: int) -> Path:
|
||||
return Path(args.output_root) / args.run_name / f"{size}x{size}"
|
||||
|
||||
|
||||
def _run_diffusers_child(args: argparse.Namespace) -> None:
|
||||
model_dir = Path(args.model_dir)
|
||||
if not model_dir.exists():
|
||||
raise RuntimeError("Set --model-dir or FLUX2_FULL_MODEL_DIR to the full Flux2 checkpoint directory")
|
||||
|
||||
output_dir = _attempt_dir(args, args.child_size)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
print(f"[image-compare] writing Diffusers image to {output_dir}", flush=True)
|
||||
_generate_diffusers_image(
|
||||
model_dir,
|
||||
output_dir / "upstream_diffusers.png",
|
||||
prompt=args.prompt,
|
||||
seed=args.seed,
|
||||
size=args.child_size,
|
||||
steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
)
|
||||
print(f"[image-compare] Diffusers success size={args.child_size}", flush=True)
|
||||
|
||||
|
||||
def _run_fastvideo_child(args: argparse.Namespace) -> None:
|
||||
model_dir = Path(args.model_dir)
|
||||
if not model_dir.exists():
|
||||
raise RuntimeError("Set --model-dir or FLUX2_FULL_MODEL_DIR to the full Flux2 checkpoint directory")
|
||||
|
||||
output_dir = _attempt_dir(args, args.child_size)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
print(f"[image-compare] writing FastVideo image to {output_dir}", flush=True)
|
||||
_generate_fastvideo_image(
|
||||
model_dir,
|
||||
output_dir / "fastvideo.png",
|
||||
prompt=args.prompt,
|
||||
seed=args.seed,
|
||||
size=args.child_size,
|
||||
steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
)
|
||||
print(f"[image-compare] FastVideo success size={args.child_size}", flush=True)
|
||||
|
||||
|
||||
def _compare_attempt(args: argparse.Namespace, size: int) -> None:
|
||||
output_dir = _attempt_dir(args, size)
|
||||
upstream = Image.open(output_dir / "upstream_diffusers.png").convert("RGB")
|
||||
fastvideo = Image.open(output_dir / "fastvideo.png").convert("RGB")
|
||||
_save_comparison(
|
||||
output_dir,
|
||||
upstream,
|
||||
fastvideo,
|
||||
{
|
||||
"prompt": args.prompt,
|
||||
"seed": args.seed,
|
||||
"size": size,
|
||||
"steps": args.steps,
|
||||
"guidance_scale": args.guidance_scale,
|
||||
"max_sequence_length": args.max_sequence_length,
|
||||
},
|
||||
)
|
||||
print(f"[image-compare] comparison success size={size} output_dir={output_dir}", flush=True)
|
||||
|
||||
|
||||
def _copy_successful_attempt(attempt_dir: Path, final_dir: Path) -> None:
|
||||
final_dir.mkdir(parents=True, exist_ok=True)
|
||||
for path in attempt_dir.iterdir():
|
||||
if path.is_file():
|
||||
shutil.copy2(path, final_dir / path.name)
|
||||
with (final_dir / "selected_attempt.txt").open("w", encoding="utf-8") as f:
|
||||
f.write(str(attempt_dir) + "\n")
|
||||
|
||||
|
||||
def _run_parent(args: argparse.Namespace) -> None:
|
||||
final_dir = Path(args.output_root) / args.run_name
|
||||
final_dir.mkdir(parents=True, exist_ok=True)
|
||||
sizes = [int(item.strip()) for item in args.sizes.split(",") if item.strip()]
|
||||
errors: list[str] = []
|
||||
for size in sizes:
|
||||
print(f"[image-compare] trying size={size}", flush=True)
|
||||
base_child_cmd = [sys.executable, __file__, *sys.argv[1:], "--child-size", str(size)]
|
||||
diffusers_result = subprocess.run([*base_child_cmd, "--child-mode", "diffusers"], check=False)
|
||||
if diffusers_result.returncode != 0:
|
||||
errors.append(f"{size}/diffusers: exit_code={diffusers_result.returncode}")
|
||||
print(
|
||||
f"[image-compare] size={size} Diffusers failed with exit_code={diffusers_result.returncode}",
|
||||
flush=True,
|
||||
)
|
||||
continue
|
||||
fastvideo_result = subprocess.run([*base_child_cmd, "--child-mode", "fastvideo"], check=False)
|
||||
if fastvideo_result.returncode == 0:
|
||||
_compare_attempt(args, size)
|
||||
attempt_dir = final_dir / f"{size}x{size}"
|
||||
_copy_successful_attempt(attempt_dir, final_dir)
|
||||
print(f"[image-compare] selected size={size}", flush=True)
|
||||
print(f"[image-compare] final output_dir={final_dir}", flush=True)
|
||||
return
|
||||
errors.append(f"{size}/fastvideo: exit_code={fastvideo_result.returncode}")
|
||||
print(
|
||||
f"[image-compare] size={size} FastVideo failed with exit_code={fastvideo_result.returncode}",
|
||||
flush=True,
|
||||
)
|
||||
raise RuntimeError("All image comparison attempts failed: " + "; ".join(errors))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = _parse_args()
|
||||
if args.child_size and args.child_mode == "diffusers":
|
||||
_run_diffusers_child(args)
|
||||
elif args.child_size and args.child_mode == "fastvideo":
|
||||
_run_fastvideo_child(args)
|
||||
else:
|
||||
_run_parent(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,865 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# pyright: reportArgumentType=false, reportAttributeAccessIssue=false, reportCallIssue=false, reportMissingTypeArgument=false
|
||||
"""Flux2 component parity tests.
|
||||
|
||||
Compares FastVideo's Flux2 components (DiT, VAE, Qwen3 and Mistral3 text
|
||||
encoders) against Diffusers/transformers references using the published
|
||||
``black-forest-labs/FLUX.2-klein-4B`` and ``black-forest-labs/FLUX.2-dev``
|
||||
weights.
|
||||
|
||||
All tests are skip-marked (CUDA + weight directory required).
|
||||
Run locally with:
|
||||
|
||||
FLUX2_MODEL_DIR=/path/to/weights pytest tests/local_tests/flux2/ -v -s
|
||||
FLUX2_FULL_MODEL_DIR=/path/to/full/weights pytest tests/local_tests/flux2/ -v -s
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from safetensors.torch import safe_open
|
||||
from torch.testing import assert_close
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="Flux2 component parity tests require CUDA",
|
||||
),
|
||||
pytest.mark.filterwarnings(
|
||||
"ignore:.*torch.jit.script_method.*:DeprecationWarning",
|
||||
),
|
||||
]
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
MODEL_DIR = Path(
|
||||
os.getenv(
|
||||
"FLUX2_MODEL_DIR",
|
||||
"/FastVideo/official_weights/black-forest-labs__FLUX.2-klein-4B",
|
||||
)
|
||||
)
|
||||
FULL_MODEL_DIR = Path(os.getenv("FLUX2_FULL_MODEL_DIR", ""))
|
||||
|
||||
|
||||
def _load_json(path: Path) -> dict:
|
||||
with path.open("r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _pick_device_and_dtype() -> tuple[torch.device, torch.dtype]:
|
||||
if torch.cuda.is_available():
|
||||
if torch.cuda.is_bf16_supported():
|
||||
return torch.device("cuda"), torch.bfloat16
|
||||
return torch.device("cuda"), torch.float32
|
||||
return torch.device("cpu"), torch.float32
|
||||
|
||||
|
||||
def _print_assert_close_means(
|
||||
label: str,
|
||||
expected: torch.Tensor,
|
||||
actual: torch.Tensor,
|
||||
) -> None:
|
||||
expected_f32 = expected.detach().float()
|
||||
actual_f32 = actual.detach().float()
|
||||
print(
|
||||
f"[{label}] assert_close means "
|
||||
f"expected_mean={expected_f32.mean().item():.6f} "
|
||||
f"actual_mean={actual_f32.mean().item():.6f} "
|
||||
f"expected_abs_mean={expected_f32.abs().mean().item():.6f} "
|
||||
f"actual_abs_mean={actual_f32.abs().mean().item():.6f}"
|
||||
)
|
||||
|
||||
|
||||
def _iter_safetensors(path: str):
|
||||
with safe_open(path, framework="pt", device="cpu") as f:
|
||||
for k in f.keys():
|
||||
yield k, f.get_tensor(k)
|
||||
|
||||
|
||||
def _iter_pretrained_safetensors(model_dir: Path):
|
||||
candidates = [
|
||||
("model.safetensors", "model.safetensors.index.json"),
|
||||
("diffusion_pytorch_model.safetensors",
|
||||
"diffusion_pytorch_model.safetensors.index.json"),
|
||||
]
|
||||
for single_name, index_name in candidates:
|
||||
single = model_dir / single_name
|
||||
if single.exists():
|
||||
yield from _iter_safetensors(str(single))
|
||||
return
|
||||
index = model_dir / index_name
|
||||
if index.exists():
|
||||
idx = _load_json(index)
|
||||
shard_names = sorted(set(idx["weight_map"].values()))
|
||||
for shard in shard_names:
|
||||
yield from _iter_safetensors(str(model_dir / shard))
|
||||
return
|
||||
|
||||
raise FileNotFoundError(
|
||||
f"Missing safetensors checkpoint in {model_dir} "
|
||||
"(expected model.safetensors or diffusion_pytorch_model.safetensors)"
|
||||
)
|
||||
|
||||
|
||||
def _require_full_model_dir() -> Path:
|
||||
if not FULL_MODEL_DIR.exists():
|
||||
pytest.skip("Set FLUX2_FULL_MODEL_DIR to activate full Flux2 component parity")
|
||||
return FULL_MODEL_DIR
|
||||
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# dist / TP fixture (single-GPU stub)
|
||||
# -----------------------------------------------------------------
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def _init_dist_and_tp_groups():
|
||||
if not torch.cuda.is_available():
|
||||
yield
|
||||
return
|
||||
|
||||
import fastvideo.distributed.parallel_state as ps
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
get_tp_group,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
|
||||
created_dist = False
|
||||
created_tp = False
|
||||
|
||||
try:
|
||||
_ = get_tp_group()
|
||||
yield
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
stubbed = False
|
||||
old_tp = old_sp = old_dp = old_world = None
|
||||
|
||||
try:
|
||||
if not dist.is_initialized():
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29500")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
os.environ.setdefault("GLOO_SOCKET_IFNAME", "lo")
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", "0")))
|
||||
|
||||
backend = "nccl" if torch.cuda.is_available() else "gloo"
|
||||
store_path = f"/tmp/fastvideo_flux2_pg_{os.getpid()}.store"
|
||||
dist.init_process_group(
|
||||
backend=backend,
|
||||
init_method=f"file://{store_path}",
|
||||
rank=int(os.environ["RANK"]),
|
||||
world_size=int(os.environ["WORLD_SIZE"]),
|
||||
)
|
||||
created_dist = True
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=int(os.environ["WORLD_SIZE"]),
|
||||
rank=int(os.environ["RANK"]),
|
||||
local_rank=int(os.environ["LOCAL_RANK"]),
|
||||
distributed_init_method="env://",
|
||||
)
|
||||
else:
|
||||
init_distributed_environment(
|
||||
world_size=dist.get_world_size(),
|
||||
rank=dist.get_rank(),
|
||||
local_rank=int(os.environ.get("LOCAL_RANK", "0")),
|
||||
distributed_init_method="env://",
|
||||
)
|
||||
|
||||
try:
|
||||
_ = get_tp_group()
|
||||
except Exception:
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
sequence_model_parallel_size=1,
|
||||
data_parallel_size=(
|
||||
dist.get_world_size() if dist.is_initialized() else 1
|
||||
),
|
||||
)
|
||||
created_tp = True
|
||||
except Exception:
|
||||
old_tp = getattr(ps, "_TP", None)
|
||||
old_sp = getattr(ps, "_SP", None)
|
||||
old_dp = getattr(ps, "_DP", None)
|
||||
old_world = getattr(ps, "_WORLD", None)
|
||||
|
||||
class _NoOpGroup:
|
||||
world_size = 1
|
||||
rank_in_group = 0
|
||||
local_rank = 0
|
||||
device_group = None
|
||||
|
||||
def all_reduce(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x
|
||||
|
||||
def all_gather(self, x: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
return x
|
||||
|
||||
def all_to_all_4D(self, x: torch.Tensor, *_a, **_kw) -> torch.Tensor:
|
||||
return x
|
||||
|
||||
def barrier(self) -> None:
|
||||
return None
|
||||
|
||||
def destroy(self) -> None:
|
||||
return None
|
||||
|
||||
ps._WORLD = _NoOpGroup()
|
||||
ps._TP = _NoOpGroup()
|
||||
ps._SP = _NoOpGroup()
|
||||
ps._DP = _NoOpGroup()
|
||||
stubbed = True
|
||||
|
||||
yield
|
||||
|
||||
if stubbed:
|
||||
ps._TP = old_tp
|
||||
ps._SP = old_sp
|
||||
ps._DP = old_dp
|
||||
ps._WORLD = old_world
|
||||
else:
|
||||
if created_tp:
|
||||
destroy_model_parallel()
|
||||
if created_dist:
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# DiT transformer parity
|
||||
# -----------------------------------------------------------------
|
||||
|
||||
def test_flux2_transformer_parity():
|
||||
"""Numerical forward parity: Diffusers Flux2Transformer2DModel vs FastVideo Flux2.
|
||||
|
||||
The two implementations expose slightly different public forward surfaces.
|
||||
This test adapts both to the same denoising-step inputs: image tokens, text
|
||||
tokens, timestep, and explicit Flux2 text/image RoPE ids for Diffusers.
|
||||
Klein does not use guidance; Diffusers' pooled projection input is provided
|
||||
as zeros because FastVideo's native Flux2 embedding intentionally ignores
|
||||
pooled text projections for this distilled path.
|
||||
"""
|
||||
transformer_dir = MODEL_DIR / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
pytest.skip(f"Flux2 transformer dir not found: {transformer_dir}")
|
||||
|
||||
from diffusers import Flux2Transformer2DModel as RefTransformer
|
||||
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
|
||||
device, dtype = _pick_device_and_dtype()
|
||||
torch.manual_seed(0)
|
||||
|
||||
cfg = _load_json(transformer_dir / "config.json")
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("Flux2Transformer2DModel")
|
||||
dit_cfg = Flux2Config()
|
||||
dit_cfg.update_model_arch(cfg)
|
||||
|
||||
in_channels = dit_cfg.in_channels
|
||||
joint_dim = dit_cfg.joint_attention_dim
|
||||
B, img_h, img_w, txt_len = 1, 8, 8, 16
|
||||
seq_len = img_h * img_w
|
||||
hidden_cpu = torch.randn(B, seq_len, in_channels, dtype=torch.float32)
|
||||
enc_cpu = torch.randn(B, txt_len, joint_dim, dtype=torch.float32)
|
||||
timestep_cpu = torch.tensor([0.5], dtype=torch.float32)
|
||||
txt_ids_cpu = torch.cartesian_prod(
|
||||
torch.arange(1), torch.arange(1), torch.arange(1), torch.arange(txt_len),
|
||||
)
|
||||
img_ids_cpu = torch.cartesian_prod(
|
||||
torch.arange(1), torch.arange(img_h), torch.arange(img_w), torch.arange(1),
|
||||
)
|
||||
|
||||
ref = RefTransformer.from_pretrained(
|
||||
str(transformer_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=False,
|
||||
).eval().to(device)
|
||||
ref.time_guidance_embed.guidance_embedder = None
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
hidden_states=hidden_cpu.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=enc_cpu.to(device=device, dtype=dtype),
|
||||
timestep=timestep_cpu.to(device=device, dtype=dtype),
|
||||
img_ids=img_ids_cpu.to(device=device),
|
||||
txt_ids=txt_ids_cpu.to(device=device),
|
||||
guidance=torch.zeros_like(timestep_cpu).to(device=device, dtype=dtype),
|
||||
return_dict=False,
|
||||
)[0].detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
fv = fv_cls(config=dit_cfg, hf_config=dict(cfg)).eval()
|
||||
fv_sd = {}
|
||||
for k, v in _iter_pretrained_safetensors(transformer_dir):
|
||||
fv_sd[k] = v
|
||||
fv.load_state_dict(fv_sd, strict=True)
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
fv_out = fv(
|
||||
hidden_states=hidden_cpu.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=enc_cpu.to(device=device, dtype=dtype),
|
||||
timestep=timestep_cpu.to(device=device, dtype=dtype),
|
||||
).detach().float().cpu()
|
||||
|
||||
assert fv_out.shape == (B, seq_len, in_channels), (
|
||||
f"Expected output shape {(B, seq_len, in_channels)}, got {fv_out.shape}"
|
||||
)
|
||||
assert torch.isfinite(fv_out).all(), "FastVideo DiT output contains non-finite values"
|
||||
assert torch.isfinite(ref_out).all(), "Diffusers DiT output contains non-finite values"
|
||||
|
||||
diff = (ref_out - fv_out).abs()
|
||||
print(
|
||||
f"[FLUX2 DIT] diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
print(
|
||||
"[FLUX2 DIT] abs-mean drift "
|
||||
f"diffusers={ref_out.abs().mean().item():.6f} "
|
||||
f"fastvideo={fv_out.abs().mean().item():.6f}"
|
||||
)
|
||||
_print_assert_close_means("FLUX2 DIT", ref_out, fv_out)
|
||||
assert_close(ref_out, fv_out, atol=1e-5, rtol=1e-5)
|
||||
|
||||
hidden_5d = hidden_cpu.reshape(B, img_h, img_w, in_channels).permute(
|
||||
0, 3, 1, 2
|
||||
).unsqueeze(2).contiguous()
|
||||
with torch.no_grad():
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
fv_out_5d = fv(
|
||||
hidden_states=hidden_5d.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=enc_cpu.to(device=device, dtype=dtype),
|
||||
timestep=timestep_cpu.to(device=device, dtype=dtype),
|
||||
).detach().float().cpu()
|
||||
fv_out_5d_seq = fv_out_5d.squeeze(2).permute(0, 2, 3, 1).reshape(
|
||||
B, seq_len, in_channels
|
||||
)
|
||||
diff_5d = (ref_out - fv_out_5d_seq).abs()
|
||||
print(
|
||||
f"[FLUX2 DIT 5D] diff max={diff_5d.max().item():.6f} "
|
||||
f"mean={diff_5d.mean().item():.6f} median={diff_5d.median().item():.6f}"
|
||||
)
|
||||
_print_assert_close_means("FLUX2 DIT 5D", ref_out, fv_out_5d_seq)
|
||||
assert_close(ref_out, fv_out_5d_seq, atol=1e-5, rtol=1e-5)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def test_flux2_full_transformer_guidance_parity():
|
||||
"""Numerical forward parity for full Flux2 transformer with embedded guidance."""
|
||||
model_dir = _require_full_model_dir()
|
||||
transformer_dir = model_dir / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
pytest.skip(f"Flux2 full transformer dir not found: {transformer_dir}")
|
||||
|
||||
from diffusers import Flux2Transformer2DModel as RefTransformer
|
||||
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
|
||||
device = torch.device(os.getenv("FLUX2_FULL_TRANSFORMER_DEVICE", "cpu"))
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
pytest.skip("FLUX2_FULL_TRANSFORMER_DEVICE=cuda requested but CUDA is unavailable")
|
||||
dtype = torch.bfloat16
|
||||
torch.manual_seed(0)
|
||||
|
||||
cfg = _load_json(transformer_dir / "config.json")
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("Flux2Transformer2DModel")
|
||||
dit_cfg = Flux2Config()
|
||||
dit_cfg.update_model_arch(cfg)
|
||||
|
||||
in_channels = dit_cfg.in_channels
|
||||
joint_dim = dit_cfg.joint_attention_dim
|
||||
B, img_h, img_w, txt_len = 1, 2, 2, 8
|
||||
seq_len = img_h * img_w
|
||||
hidden_cpu = torch.randn(B, seq_len, in_channels, dtype=torch.float32)
|
||||
enc_cpu = torch.randn(B, txt_len, joint_dim, dtype=torch.float32)
|
||||
timestep_cpu = torch.tensor([0.5], dtype=torch.float32)
|
||||
guidance_cpu = torch.tensor([4.0], dtype=torch.float32)
|
||||
txt_ids_cpu = torch.cartesian_prod(
|
||||
torch.arange(1), torch.arange(1), torch.arange(1), torch.arange(txt_len),
|
||||
)
|
||||
img_ids_cpu = torch.cartesian_prod(
|
||||
torch.arange(1), torch.arange(img_h), torch.arange(img_w), torch.arange(1),
|
||||
)
|
||||
|
||||
ref = RefTransformer.from_pretrained(
|
||||
str(transformer_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=False,
|
||||
).eval()
|
||||
if device.type != "cpu":
|
||||
ref = ref.to(device)
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
hidden_states=hidden_cpu.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=enc_cpu.to(device=device, dtype=dtype),
|
||||
timestep=timestep_cpu.to(device=device, dtype=dtype),
|
||||
img_ids=img_ids_cpu.to(device=device),
|
||||
txt_ids=txt_ids_cpu.to(device=device),
|
||||
guidance=guidance_cpu.to(device=device, dtype=dtype),
|
||||
return_dict=False,
|
||||
)[0].detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
old_default_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(dtype)
|
||||
try:
|
||||
fv = fv_cls(config=dit_cfg, hf_config=dict(cfg)).eval()
|
||||
finally:
|
||||
torch.set_default_dtype(old_default_dtype)
|
||||
fv_sd = {}
|
||||
for k, v in _iter_pretrained_safetensors(transformer_dir):
|
||||
fv_sd[k] = v
|
||||
fv.load_state_dict(fv_sd, strict=True)
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
fv_out = fv(
|
||||
hidden_states=hidden_cpu.to(device=device, dtype=dtype),
|
||||
encoder_hidden_states=enc_cpu.to(device=device, dtype=dtype),
|
||||
timestep=timestep_cpu.to(device=device, dtype=dtype),
|
||||
guidance=guidance_cpu.to(device=device, dtype=dtype),
|
||||
img_ids=img_ids_cpu.to(device=device),
|
||||
txt_ids=txt_ids_cpu.to(device=device),
|
||||
).detach().float().cpu()
|
||||
|
||||
assert fv_out.shape == (B, seq_len, in_channels), (
|
||||
f"Expected output shape {(B, seq_len, in_channels)}, got {fv_out.shape}"
|
||||
)
|
||||
assert torch.isfinite(fv_out).all(), "FastVideo full DiT output contains non-finite values"
|
||||
assert torch.isfinite(ref_out).all(), "Diffusers full DiT output contains non-finite values"
|
||||
|
||||
diff = (ref_out - fv_out).abs()
|
||||
print(
|
||||
f"[FLUX2 FULL DIT] diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
_print_assert_close_means("FLUX2 FULL DIT", ref_out, fv_out)
|
||||
assert_close(ref_out, fv_out, atol=1e-5, rtol=1e-5)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# VAE parity
|
||||
# -----------------------------------------------------------------
|
||||
|
||||
def test_flux2_vae_encode_decode_parity():
|
||||
"""Encode/decode parity: Diffusers AutoencoderKLFlux2 vs FastVideo Flux2 VAE."""
|
||||
vae_dir = MODEL_DIR / "vae"
|
||||
if not vae_dir.exists():
|
||||
pytest.skip(f"Flux2 VAE dir not found: {vae_dir}")
|
||||
|
||||
from diffusers import AutoencoderKLFlux2 as RefVAE
|
||||
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
|
||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
dtype = torch.float32
|
||||
torch.manual_seed(0)
|
||||
|
||||
ref = RefVAE.from_pretrained(
|
||||
str(vae_dir), local_files_only=True, torch_dtype=dtype,
|
||||
).eval().to(device)
|
||||
|
||||
x = torch.randn(1, 3, 64, 64, device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_latents = ref.encode(x).latent_dist.mean.detach().float().cpu()
|
||||
ref_dec = ref.decode(
|
||||
ref_latents.to(device=device, dtype=dtype)
|
||||
).sample.detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
cfg = _load_json(vae_dir / "config.json")
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("AutoencoderKLFlux2")
|
||||
vae_cfg = Flux2VAEConfig()
|
||||
vae_cfg.update_model_arch(cfg)
|
||||
fv = fv_cls(vae_cfg).eval()
|
||||
|
||||
weight_path = vae_dir / "diffusion_pytorch_model.safetensors"
|
||||
fv_sd = safetensors_load_file(str(weight_path), device="cpu")
|
||||
fv.load_state_dict(fv_sd, strict=True)
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_latents = fv.encode(x).mean.detach().float().cpu()
|
||||
fv_dec_output = fv.decode(
|
||||
fv_latents.to(device=device, dtype=dtype)
|
||||
)
|
||||
fv_dec_sample = getattr(fv_dec_output, "sample", fv_dec_output)
|
||||
fv_dec = fv_dec_sample.detach().float().cpu()
|
||||
|
||||
_print_assert_close_means("FLUX2 VAE encode", ref_latents, fv_latents)
|
||||
assert_close(ref_latents, fv_latents, atol=1e-4, rtol=1e-4)
|
||||
_print_assert_close_means("FLUX2 VAE decode", ref_dec, fv_dec)
|
||||
assert_close(ref_dec, fv_dec, atol=1e-4, rtol=1e-4)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def test_flux2_full_vae_encode_decode_parity():
|
||||
"""Encode/decode parity for the full Flux2 VAE config and weights."""
|
||||
model_dir = _require_full_model_dir()
|
||||
vae_dir = model_dir / "vae"
|
||||
if not vae_dir.exists():
|
||||
pytest.skip(f"Flux2 full VAE dir not found: {vae_dir}")
|
||||
|
||||
from diffusers import AutoencoderKLFlux2 as RefVAE
|
||||
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
|
||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
dtype = torch.float32
|
||||
torch.manual_seed(0)
|
||||
|
||||
ref = RefVAE.from_pretrained(
|
||||
str(vae_dir), local_files_only=True, torch_dtype=dtype,
|
||||
).eval().to(device)
|
||||
|
||||
x = torch.randn(1, 3, 64, 64, device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_latents = ref.encode(x).latent_dist.mean.detach().float().cpu()
|
||||
ref_dec = ref.decode(
|
||||
ref_latents.to(device=device, dtype=dtype)
|
||||
).sample.detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
cfg = _load_json(vae_dir / "config.json")
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("AutoencoderKLFlux2")
|
||||
vae_cfg = Flux2VAEConfig()
|
||||
vae_cfg.update_model_arch(cfg)
|
||||
fv = fv_cls(vae_cfg).eval()
|
||||
|
||||
weight_path = vae_dir / "diffusion_pytorch_model.safetensors"
|
||||
fv_sd = safetensors_load_file(str(weight_path), device="cpu")
|
||||
fv.load_state_dict(fv_sd, strict=True)
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_latents = fv.encode(x).mean.detach().float().cpu()
|
||||
fv_dec_output = fv.decode(fv_latents.to(device=device, dtype=dtype))
|
||||
fv_dec_sample = getattr(fv_dec_output, "sample", fv_dec_output)
|
||||
fv_dec = fv_dec_sample.detach().float().cpu()
|
||||
|
||||
_print_assert_close_means("FLUX2 FULL VAE encode", ref_latents, fv_latents)
|
||||
assert_close(ref_latents, fv_latents, atol=1e-4, rtol=1e-4)
|
||||
_print_assert_close_means("FLUX2 FULL VAE decode", ref_dec, fv_dec)
|
||||
assert_close(ref_dec, fv_dec, atol=1e-4, rtol=1e-4)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# Qwen3 text encoder parity
|
||||
# -----------------------------------------------------------------
|
||||
|
||||
def test_flux2_qwen3_text_encoder_parity():
|
||||
"""Hidden-state parity for the Flux2 Qwen3 loader path.
|
||||
|
||||
Flux2 uses the HuggingFace Qwen3 module through FastVideo's component
|
||||
loader, so this validates that passthrough path against a direct HF load.
|
||||
The native TP-aware Qwen3 class is intentionally not used by the Flux2
|
||||
pipeline until it can provide strict hidden-state parity.
|
||||
"""
|
||||
text_encoder_dir = MODEL_DIR / "text_encoder"
|
||||
if not text_encoder_dir.exists():
|
||||
pytest.skip(f"Flux2 text_encoder dir not found: {text_encoder_dir}")
|
||||
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||
|
||||
device, dtype = _pick_device_and_dtype()
|
||||
torch.manual_seed(0)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
str(MODEL_DIR / "tokenizer"), local_files_only=True,
|
||||
)
|
||||
prompt = "a photo of a cat"
|
||||
formatted = tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": prompt}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
toks = tokenizer(
|
||||
[formatted],
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = toks["input_ids"].to(device=device)
|
||||
attention_mask = toks["attention_mask"].to(device=device)
|
||||
|
||||
ref = AutoModelForCausalLM.from_pretrained(
|
||||
str(text_encoder_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval().to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
ref_embeds = torch.stack(
|
||||
[ref_out.hidden_states[k] for k in (9, 18, 27)], dim=1
|
||||
)
|
||||
ref_embeds = ref_embeds.permute(0, 2, 1, 3).reshape(
|
||||
input_ids.shape[0],
|
||||
input_ids.shape[1],
|
||||
-1,
|
||||
).detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.models.encoders.qwen3 import Qwen3ForCausalLM
|
||||
|
||||
cfg_raw = _load_json(text_encoder_dir / "config.json")
|
||||
for k in ("_name_or_path", "transformers_version", "model_type", "torch_dtype"):
|
||||
cfg_raw.pop(k, None)
|
||||
|
||||
fv_cfg = Qwen3TextConfig()
|
||||
fv_cfg.update_model_arch(cfg_raw)
|
||||
fv = Qwen3ForCausalLM.from_pretrained_local(
|
||||
str(text_encoder_dir),
|
||||
fv_cfg,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
assert fv.__class__.__module__.startswith("transformers"), (
|
||||
"Flux2 Qwen3 should load through the exact HuggingFace passthrough path"
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_out = fv(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
fv_embeds = torch.stack(
|
||||
[fv_out.hidden_states[k] for k in (9, 18, 27)], dim=1
|
||||
)
|
||||
fv_embeds = fv_embeds.permute(0, 2, 1, 3).reshape(
|
||||
input_ids.shape[0],
|
||||
input_ids.shape[1],
|
||||
-1,
|
||||
).detach().float().cpu()
|
||||
|
||||
diff = (ref_embeds - fv_embeds).abs()
|
||||
print(
|
||||
f"[FLUX2 QWEN3] diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
_print_assert_close_means("FLUX2 QWEN3", ref_embeds, fv_embeds)
|
||||
assert_close(ref_embeds, fv_embeds, atol=1e-5, rtol=1e-5)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def test_flux2_mistral3_text_encoder_parity():
|
||||
"""Hidden-state parity for the full Flux2 Mistral3 HF passthrough path."""
|
||||
model_dir = _require_full_model_dir()
|
||||
text_encoder_dir = model_dir / "text_encoder"
|
||||
tokenizer_dir = model_dir / "tokenizer"
|
||||
if not text_encoder_dir.exists():
|
||||
pytest.skip(f"Flux2 full text_encoder dir not found: {text_encoder_dir}")
|
||||
if not tokenizer_dir.exists():
|
||||
pytest.skip(f"Flux2 full tokenizer dir not found: {tokenizer_dir}")
|
||||
|
||||
from transformers import AutoModelForImageTextToText, AutoProcessor
|
||||
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.models.encoders.mistral3 import Mistral3ForConditionalGeneration
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_text_encoding import (
|
||||
FLUX2_SYSTEM_MESSAGE,
|
||||
_format_flux2_full_input,
|
||||
)
|
||||
|
||||
device = torch.device(os.getenv("FLUX2_MISTRAL3_DEVICE", "cpu"))
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
pytest.skip("FLUX2_MISTRAL3_DEVICE=cuda requested but CUDA is unavailable")
|
||||
dtype = torch.bfloat16
|
||||
torch.manual_seed(0)
|
||||
|
||||
processor = AutoProcessor.from_pretrained(str(tokenizer_dir), local_files_only=True)
|
||||
messages = _format_flux2_full_input(["a photo of a cat"], FLUX2_SYSTEM_MESSAGE)
|
||||
toks = processor.apply_chat_template(
|
||||
messages,
|
||||
add_generation_prompt=False,
|
||||
tokenize=True,
|
||||
return_dict=True,
|
||||
return_tensors="pt",
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=64,
|
||||
)
|
||||
input_ids = toks["input_ids"].to(device=device)
|
||||
attention_mask = toks["attention_mask"].to(device=device)
|
||||
|
||||
ref = AutoModelForImageTextToText.from_pretrained(
|
||||
str(text_encoder_dir),
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval()
|
||||
if device.type != "cpu":
|
||||
ref = ref.to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
use_cache=False,
|
||||
)
|
||||
ref_embeds = torch.stack(
|
||||
[ref_out.hidden_states[k] for k in (10, 20, 30)], dim=1
|
||||
)
|
||||
ref_embeds = ref_embeds.permute(0, 2, 1, 3).reshape(
|
||||
input_ids.shape[0],
|
||||
input_ids.shape[1],
|
||||
-1,
|
||||
).detach().float().cpu()
|
||||
|
||||
del ref
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
cfg_raw = _load_json(text_encoder_dir / "config.json")
|
||||
for k in ("_name_or_path", "transformers_version", "model_type", "torch_dtype"):
|
||||
cfg_raw.pop(k, None)
|
||||
|
||||
fv_cfg = Mistral3TextConfig()
|
||||
fv_cfg.update_model_arch(cfg_raw)
|
||||
fv = Mistral3ForConditionalGeneration.from_pretrained_local(
|
||||
str(text_encoder_dir),
|
||||
fv_cfg,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
assert fv.__class__.__module__.startswith("transformers"), (
|
||||
"Flux2 Mistral3 should load through the exact HuggingFace passthrough path"
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_out = fv(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
use_cache=False,
|
||||
)
|
||||
fv_embeds = torch.stack(
|
||||
[fv_out.hidden_states[k] for k in (10, 20, 30)], dim=1
|
||||
)
|
||||
fv_embeds = fv_embeds.permute(0, 2, 1, 3).reshape(
|
||||
input_ids.shape[0],
|
||||
input_ids.shape[1],
|
||||
-1,
|
||||
).detach().float().cpu()
|
||||
|
||||
diff = (ref_embeds - fv_embeds).abs()
|
||||
print(
|
||||
f"[FLUX2 MISTRAL3] diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
_print_assert_close_means("FLUX2 MISTRAL3", ref_embeds, fv_embeds)
|
||||
assert_close(ref_embeds, fv_embeds, atol=1e-5, rtol=1e-5)
|
||||
|
||||
del fv
|
||||
gc.collect()
|
||||
if device.type == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,332 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Smoke / preflight tests for the Flux2 Klein T2I pipeline.
|
||||
|
||||
The preflight validates import, registry, preset, and config wiring in an
|
||||
environment with the expected optional packages. The load/generate smoke is
|
||||
activated with CUDA plus
|
||||
``FLUX2_MODEL_DIR=/path/to/black-forest-labs__FLUX.2-klein-4B``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
MODEL_DIR = Path(os.getenv("FLUX2_MODEL_DIR", ""))
|
||||
FULL_MODEL_DIR = Path(os.getenv("FLUX2_FULL_MODEL_DIR", ""))
|
||||
FULL_HEIGHT = int(os.getenv("FLUX2_FULL_HEIGHT", "128"))
|
||||
FULL_WIDTH = int(os.getenv("FLUX2_FULL_WIDTH", "128"))
|
||||
FULL_NUM_INFERENCE_STEPS = int(os.getenv("FLUX2_FULL_STEPS", "1"))
|
||||
FULL_GUIDANCE_SCALE = float(os.getenv("FLUX2_FULL_GUIDANCE_SCALE", "4.0"))
|
||||
FULL_MAX_SEQUENCE_LENGTH = int(os.getenv("FLUX2_FULL_MAX_SEQUENCE_LENGTH", "64"))
|
||||
FULL_NUM_GPUS = int(os.getenv("FLUX2_FULL_NUM_GPUS", "2"))
|
||||
FULL_TP_SIZE = int(os.getenv("FLUX2_FULL_TP_SIZE", str(FULL_NUM_GPUS)))
|
||||
FULL_SP_SIZE = int(
|
||||
os.getenv(
|
||||
"FLUX2_FULL_SP_SIZE",
|
||||
"1" if FULL_NUM_GPUS > 1 else str(FULL_NUM_GPUS),
|
||||
)
|
||||
)
|
||||
requires_flux2_runtime = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="Flux2 pipeline imports require the CUDA/kernel runtime",
|
||||
)
|
||||
|
||||
|
||||
@requires_flux2_runtime
|
||||
def test_flux2_full_typed_surface_preflight() -> None:
|
||||
"""Import + registry + preset wiring preflight for full Flux2."""
|
||||
import fastvideo.registry as registry
|
||||
from fastvideo.api.presets import get_preset, get_presets_for_family
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.pipelines.flux_2 import Flux2PipelineConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_pipeline import (
|
||||
EntryClass,
|
||||
Flux2Pipeline,
|
||||
)
|
||||
|
||||
assert Flux2Pipeline.__name__ == "Flux2Pipeline"
|
||||
assert EntryClass is Flux2Pipeline
|
||||
|
||||
default_preset, model_family = registry.get_preset_selection(
|
||||
"black-forest-labs/FLUX.2-dev"
|
||||
)
|
||||
assert model_family == "flux2"
|
||||
assert default_preset == "flux2_dev"
|
||||
|
||||
info = registry.get_model_info(
|
||||
"black-forest-labs/FLUX.2-dev",
|
||||
workload_type=WorkloadType.T2I,
|
||||
override_pipeline_cls_name="Flux2Pipeline",
|
||||
)
|
||||
assert info.pipeline_cls is Flux2Pipeline
|
||||
assert info.pipeline_config_cls is Flux2PipelineConfig
|
||||
|
||||
names = {p.name for p in get_presets_for_family("flux2")}
|
||||
assert "flux2_dev" in names
|
||||
preset = get_preset("flux2_dev", "flux2")
|
||||
assert preset.defaults["num_inference_steps"] == 50
|
||||
assert preset.defaults["height"] == 1024
|
||||
assert preset.defaults["width"] == 1024
|
||||
assert preset.defaults["guidance_scale"] == 4.0
|
||||
assert preset.defaults["num_frames"] == 1
|
||||
|
||||
cfg = Flux2PipelineConfig()
|
||||
assert cfg.embedded_cfg_scale == 4.0
|
||||
assert cfg.flux2_text_encoder_type == "mistral3"
|
||||
assert cfg.text_encoder_out_layers == (10, 20, 30)
|
||||
assert isinstance(cfg.text_encoder_configs[0], Mistral3TextConfig)
|
||||
model_cls, arch = ModelRegistry.resolve_model_cls(
|
||||
"Mistral3ForConditionalGeneration"
|
||||
)
|
||||
assert arch == "Mistral3ForConditionalGeneration"
|
||||
assert model_cls.__name__ == "Mistral3ForConditionalGeneration"
|
||||
|
||||
|
||||
class _FakeFlux2Processor:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[list[list[dict[str, Any]]], dict[str, Any]]] = []
|
||||
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages: list[list[dict[str, Any]]],
|
||||
**kwargs: Any,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
self.calls.append((messages, kwargs))
|
||||
assert kwargs["add_generation_prompt"] is False
|
||||
assert kwargs["tokenize"] is True
|
||||
assert kwargs["padding"] == "max_length"
|
||||
assert kwargs["truncation"] is True
|
||||
assert messages[0][0]["role"] == "system"
|
||||
assert messages[0][1]["role"] == "user"
|
||||
batch_size = len(messages)
|
||||
max_length = int(kwargs["max_length"])
|
||||
input_ids = torch.arange(max_length, dtype=torch.long).repeat(
|
||||
batch_size,
|
||||
1,
|
||||
)
|
||||
attention_mask = torch.ones(batch_size, max_length, dtype=torch.long)
|
||||
return {"input_ids": input_ids, "attention_mask": attention_mask}
|
||||
|
||||
|
||||
class _FakeMistral3Encoder(nn.Module):
|
||||
|
||||
def __init__(self, hidden_size: int = 4) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.zeros(1))
|
||||
self._hidden_size = hidden_size
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
return self.weight.dtype
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
output_hidden_states: bool,
|
||||
use_cache: bool,
|
||||
**_kwargs: Any,
|
||||
) -> Any:
|
||||
assert output_hidden_states is True
|
||||
assert use_cache is False
|
||||
assert attention_mask.shape == input_ids.shape
|
||||
base = input_ids.to(self.weight.dtype).unsqueeze(-1).expand(
|
||||
*input_ids.shape,
|
||||
self._hidden_size,
|
||||
)
|
||||
hidden_states = tuple(base + float(i) for i in range(31))
|
||||
return SimpleNamespace(hidden_states=hidden_states)
|
||||
|
||||
|
||||
@requires_flux2_runtime
|
||||
def test_flux2_full_text_stage_uses_mistral3_format_and_embedded_guidance() -> None:
|
||||
"""Full Flux2 text encoding uses Mistral3 formatting and disables generic CFG."""
|
||||
from fastvideo.configs.pipelines.flux_2 import Flux2PipelineConfig
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_text_encoding import (
|
||||
Flux2TextEncodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
processor = _FakeFlux2Processor()
|
||||
encoder = _FakeMistral3Encoder()
|
||||
stage = Flux2TextEncodingStage(text_encoders=[encoder], tokenizers=[processor])
|
||||
cfg = Flux2PipelineConfig()
|
||||
cfg.text_encoder_out_layers = (10, 20, 30)
|
||||
args = SimpleNamespace(pipeline_config=cfg)
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="image",
|
||||
prompt="a cat [IMG] on a chair",
|
||||
guidance_scale=4.0,
|
||||
negative_prompt="should not be encoded",
|
||||
)
|
||||
assert batch.do_classifier_free_guidance is True
|
||||
|
||||
out = stage.forward(batch, cast(Any, args))
|
||||
|
||||
assert out.do_classifier_free_guidance is False
|
||||
assert out.negative_prompt_embeds == []
|
||||
assert len(out.prompt_embeds) == 1
|
||||
assert out.prompt_embeds[0].shape == (1, 512, 12)
|
||||
assert out.extra["flux2_txt_ids"].shape == (1, 512, 4)
|
||||
assert out.extra["flux2_txt_ids"][0, -1].tolist() == [0, 0, 0, 511]
|
||||
assert len(processor.calls) == 1
|
||||
messages, _kwargs = processor.calls[0]
|
||||
assert "[IMG]" not in messages[0][1]["content"][0]["text"]
|
||||
|
||||
|
||||
@requires_flux2_runtime
|
||||
def test_flux2_klein_typed_surface_preflight() -> None:
|
||||
"""Import + registry + preset wiring preflight."""
|
||||
import fastvideo.registry as registry
|
||||
from fastvideo.api.presets import get_preset, get_presets_for_family
|
||||
from fastvideo.configs.pipelines.flux_2 import Flux2KleinPipelineConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.pipelines.basic.flux_2.flux_2_klein_pipeline import (
|
||||
EntryClass,
|
||||
Flux2KleinPipeline,
|
||||
)
|
||||
|
||||
assert Flux2KleinPipeline.__name__ == "Flux2KleinPipeline"
|
||||
assert EntryClass is Flux2KleinPipeline
|
||||
|
||||
default_preset, model_family = registry.get_preset_selection(
|
||||
"black-forest-labs/FLUX.2-klein-4B"
|
||||
)
|
||||
assert model_family == "flux2"
|
||||
assert default_preset == "flux2_klein_4b"
|
||||
|
||||
info = registry.get_model_info(
|
||||
"black-forest-labs/FLUX.2-klein-4B",
|
||||
workload_type=WorkloadType.T2I,
|
||||
override_pipeline_cls_name="Flux2KleinPipeline",
|
||||
)
|
||||
assert info.pipeline_cls is Flux2KleinPipeline
|
||||
assert info.pipeline_config_cls is Flux2KleinPipelineConfig
|
||||
|
||||
names = {p.name for p in get_presets_for_family("flux2")}
|
||||
assert "flux2_klein_4b" in names
|
||||
preset = get_preset("flux2_klein_4b", "flux2")
|
||||
assert preset.defaults["num_inference_steps"] == 4
|
||||
assert preset.defaults["height"] == 1024
|
||||
assert preset.defaults["width"] == 1024
|
||||
assert preset.defaults["guidance_scale"] == 1.0
|
||||
assert preset.defaults["num_frames"] == 1
|
||||
|
||||
assert Flux2KleinPipeline._required_config_modules == [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="Flux2 Klein pipeline load/generate smoke requires CUDA",
|
||||
)
|
||||
def test_flux2_klein_pipeline_load_generate_smoke() -> None:
|
||||
"""Optional real load + four-step latent generate smoke for local weights."""
|
||||
if not MODEL_DIR.exists():
|
||||
pytest.skip("Set FLUX2_MODEL_DIR to activate Flux2 Klein load/generate smoke")
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(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="Flux2KleinPipeline",
|
||||
)
|
||||
try:
|
||||
result = generator.generate_video(
|
||||
prompt="a photo of a banana on a wooden table, studio lighting",
|
||||
output_path="outputs_video/flux2_klein_smoke",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=1024,
|
||||
width=1024,
|
||||
num_frames=1,
|
||||
num_inference_steps=4,
|
||||
guidance_scale=1.0,
|
||||
seed=0,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
assert isinstance(result, dict)
|
||||
result_dict = cast(dict[str, Any], result)
|
||||
samples = result_dict["samples"]
|
||||
assert torch.is_tensor(samples)
|
||||
assert samples.ndim in (3, 5)
|
||||
assert torch.isfinite(samples).all()
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="Flux2 full pipeline load/generate smoke requires CUDA",
|
||||
)
|
||||
def test_flux2_full_pipeline_load_generate_smoke() -> None:
|
||||
"""Optional real load + short latent generate smoke for full Flux2 weights."""
|
||||
if not FULL_MODEL_DIR.exists():
|
||||
pytest.skip("Set FLUX2_FULL_MODEL_DIR to activate Flux2 full load/generate smoke")
|
||||
if torch.cuda.device_count() < FULL_NUM_GPUS:
|
||||
pytest.skip(
|
||||
f"Flux2 full load/generate smoke requires {FULL_NUM_GPUS} CUDA devices; "
|
||||
f"found {torch.cuda.device_count()}"
|
||||
)
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(FULL_MODEL_DIR),
|
||||
num_gpus=FULL_NUM_GPUS,
|
||||
tp_size=FULL_TP_SIZE,
|
||||
sp_size=FULL_SP_SIZE,
|
||||
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="Flux2Pipeline",
|
||||
)
|
||||
try:
|
||||
result = generator.generate_video(
|
||||
prompt="a photo of a banana on a wooden table, studio lighting",
|
||||
output_path="outputs_video/flux2_full_smoke",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=FULL_HEIGHT,
|
||||
width=FULL_WIDTH,
|
||||
num_frames=1,
|
||||
num_inference_steps=FULL_NUM_INFERENCE_STEPS,
|
||||
guidance_scale=FULL_GUIDANCE_SCALE,
|
||||
max_sequence_length=FULL_MAX_SEQUENCE_LENGTH,
|
||||
seed=0,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
assert isinstance(result, dict)
|
||||
result_dict = cast(dict[str, Any], result)
|
||||
samples = result_dict["samples"]
|
||||
assert torch.is_tensor(samples)
|
||||
assert samples.ndim in (3, 5)
|
||||
assert torch.isfinite(samples).all()
|
||||
Reference in New Issue
Block a user