Compare commits

..
Author SHA1 Message Date
SolitaryThinker d1a1b52f78 [docs]: add release skill for future version bumps 2026-06-04 14:29:35 -07:00
140 changed files with 525 additions and 19936 deletions
-32
View File
@@ -1,32 +0,0 @@
---
name: add-reward-model
description: Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.
---
# Add Reward Model
Use for reward models consumed by RL methods.
## Placement
- Put reusable reward code under `fastvideo/train/methods/rl/rewards/`.
- Expose public builders from `fastvideo/train/methods/rl/rewards/__init__.py`.
- Keep method-specific aggregation or advantage logic out of reward classes.
## Media Inputs
- Reward callables receive decoded media tensors.
- Accept single-frame tensors as `[B, C, H, W]` and multi-frame tensors as `[B, C, T, H, W]` when practical.
- Frame selection is reward-specific. Frame scorers such as PickScore and CLIPScore should explicitly select frame `0`; temporal rewards should inspect whichever frames they need.
- Return one scalar reward per prompt/sample.
## Attribution
- If code is ported or closely adapted from another repo, add a short comment or docstring naming the source file/function.
- Preserve SPDX headers used by FastVideo files.
## Tests
- Unit-test tensor layout handling without loading large reward checkpoints.
- Allow fake scorer injection for multi-reward tests.
- Test weighted reward aggregation and metric keys.
-38
View File
@@ -1,38 +0,0 @@
---
name: add-rl-method
description: Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.
---
# Add RL Method
Use for new RL methods in the modular `fastvideo/train` stack.
## Required Shape
- Add the method under `fastvideo/train/methods/rl/`.
- Subclass `TrainingMethod`.
- Keep model-family logic in `ModelBase` wrappers.
- Decode generated latents through `ModelBase.decode_latents`; add that hook to the new model wrapper instead of decoding inside the RL method.
- Use `fastvideo/train/methods/rl/common/sampling.py` for generation unless the method has a documented reason to avoid sampling.
- Use `fastvideo/train/methods/rl/common/prompt_sampling.py` for reusable grouped prompt sampling patterns such as DiffusionNFT K-repeat.
- Use `fastvideo/train/methods/rl/rewards/` for reward models.
## Optimization
- Return `manages_optimization() == True` only when the method must own a nonstandard outer/inner loop.
- If using managed optimization, implement `managed_train_step(data_stream, iteration)`.
- Existing trainer callbacks, checkpointing, tracking, and validation should still work.
## Config
- Put method knobs under `method`.
- Put sampler knobs under `method.sampling`.
- Do not put scheduler or trajectory policy into model configs.
- Do not split a diffusers-style scheduler from its built-in `step()` solver in YAML; use `trajectory` only for higher-level ODE vs re-noise behavior.
- Avoid fixed timestep lists in examples unless reproducing a known baseline; prefer scheduler-generated defaults.
## Tests
- Add fake-model tests for sampler/method behavior.
- Add config parse tests for the public YAML.
- Confirm existing train methods stay on the default Trainer path.
+1 -3
View File
@@ -8,6 +8,4 @@
{"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": "rlhf-training-abstractions", "description": "Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.", "path": "rlhf-training-abstractions/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-rl-method", "description": "Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.", "path": "add-rl-method/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-reward-model", "description": "Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.", "path": "add-reward-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"}
+248
View File
@@ -0,0 +1,248 @@
---
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.
@@ -1,41 +0,0 @@
---
name: rlhf-training-abstractions
description: Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.
---
# RLHF Training Abstractions
Use this skill before editing RLHF-style training code in `fastvideo/train`.
## Boundaries
- RL methods live under `fastvideo/train/methods/rl/` and own algorithm logic: reward collection, advantage computation, policy loss, KL/reference terms, and optimizer cadence.
- Rewards live under `fastvideo/train/methods/rl/rewards/` and must be reusable across RL methods.
- RL methods pass decoded media to rewards; each reward decides whether to use the first frame, sampled frames, or the full video.
- Sampling lives under `fastvideo/train/methods/rl/common/` and must use `ModelBase` primitives plus scheduler math, not model-family inference pipelines.
- Model wrappers under `fastvideo/train/models/` own model-specific forward details.
- Model wrappers also own model-specific latent decoding via `ModelBase.decode_latents`; RL methods should not reach into VAE normalization internals.
- Shared RL helpers such as K-repeat prompt sampling belong under `fastvideo/train/methods/rl/common/` when they are reusable across RL methods.
## Anti-Patterns
- Do not bind RL methods to inference pipeline classes such as `WanDMDPipeline`.
- Do not hardcode timestep lists in a method when the scheduler can generate them.
- Do not put reward-model code inside one RL method.
- Do not make existing non-RL methods use method-managed optimization unless explicitly requested.
## Sampling Policy
- Prefer YAML-configured `method.sampling` with `scheduler`, `trajectory`, `num_steps`, `timesteps`, and `sigmas`.
- Treat diffusers-style scheduler classes as owning both the timestep schedule and their `step()` update rule; avoid a separate `solver` field unless a new sampler truly implements solver math outside the scheduler object.
- Missing `timesteps` means “ask the scheduler”; explicit `timesteps` or `sigmas` are overrides.
- ODE-style trajectories should not re-noise between denoising steps.
- SDE/re-noise behavior must be explicit in config.
## Validation
- Run focused local tests for sampler config and Trainer opt-in behavior.
- Verify existing train methods still report `manages_optimization() == False`.
- Keep fixed-prompt validation helpers in `fastvideo/train/methods/rl/common/validation.py` so new RL methods can reuse sharding and captions.
- Test distributed prompt grouping helpers separately from heavyweight model loading.
- Run `pre-commit run --files <changed paths>`; respect configured excludes.
-2
View File
@@ -34,8 +34,6 @@ env
*.log
weights/
logs/
official_weights/
converted_weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
-2294
View File
File diff suppressed because it is too large Load Diff
-71
View File
@@ -1,71 +0,0 @@
# FastVideo Next-Gen Runtime — Two-Page Summary
**Companion to** `design.md` (v19, 2026-06-12) · **Status:** draft for discussion · **Ask:** read this, then dive into the sections you own.
---
## The problem
FastVideo's pipeline abstraction has been outgrown by its own model zoo. Four facts, all on `main` today:
- **The denoise/sampling loop exists in four copies** — inference stages (`pipelines/stages/denoising.py`, a 1,381-line file), `train/` distillation methods, legacy `training/` monoliths, and the just-landed RL work. The fourth copy documents the cause in its own docstring: `DiffusionSampler` (PR #1450) *"intentionally does not call FastVideo's full inference pipelines"* because the only consumable units are family-bound pipeline classes. Every new post-training method must pick between a wrong dependency and a private loop.
- **There is no serving runtime.** One request at a time, no queue, no cross-request batching. Dreamverse — our shipping product — hand-rolls its own GPU pool, queue, warmup, and streaming relay at a cost of **one B200 per user session**.
- **Cosmos3 outgrows both the stage abstraction and the alternatives.** Its AR text reasoner and multimodal diffusion denoiser are *the same resident weights* driven by two loop types within one request — our full port (incl. action modality, on `feat/cosmos3-reasoning`) runs it only by bypassing stages with one monolithic block. Multi-engine DAG stacks (vllm-omni, sglang-omni) compose separable stages with disjoint weights; none can express a mixture-of-transformers model. This is the forcing function.
- **RL has arrived and pays the tax already** — likelihood-free DiffusionNFT for Wan (#1450): in-process rollouts with zero serving-grade optimizations (no CFG, dense attention, full 25-step ODE), a vendored sampler, and a parallel validation path built *because* inference pipelines aren't consumable as a library.
## The design
```
Request plane OpenAI-compatible server (videos/images/audio/chat) · AsyncEngine
OmniRequest ─► admission ─► queue ─► OmniOutput (typed modality parts)
Pipeline plane PipelineSpec: declarative graph per family
nodes: Stage | LoopStage (DenoiseLoop, ARDecodeLoop) · typed Artifacts
policies: CFG, ExpertRouting, AttnMetadata, Precision, FlowShift
Execution plane StepScheduler (multiplexes denoise + AR steps) · worker pools (TP/SP + CFG-parallel)
CacheManager (paged text-KV ▪ chunked causal-video KV ▪ feature caches) · connectors
```
Default deployment is exactly today's: one SPMD pool, co-located nodes, synchronous call. Serving is additive configuration, not a different code path. Five load-bearing decisions:
1. **Loop inversion.** Loops become `LoopStage` nodes exposing `init / step / finalize`; the runtime owns iteration, families own step bodies (with a custom-step escape hatch — the runtime never dictates step factoring). This is the enabler for step-level scheduling, streaming, MoT interleaving, and one shared loop across inference, distillation, and RL. Stated honestly: **no surveyed system does this at scheduler granularity** — the risk is retired by Phase-1 bit-identical parity gates and a measured falsifier, not borrowed validation.
2. **Cost-model scheduling.** Denoise steps and AR tokens are incommensurable (bidirectional attention is O(L²) per step with zero KV amortization; steps differ ~1000×) — the budget currency is **predicted GPU-time** from a per-(model, phase, shape) cost model, calibrated by the profiler and published to Dynamo's router/Planner as the same artifact.
3. **One substrate for inference, training, and RL** — models, loaders, configs, schedulers, parallel state, loop step bodies — under a strict `engine never imports train` rule. Trainers keep their internals; their embedded sampling paths migrate onto the shared loops.
4. **Consistency is a declared, measured contract**, not a hope: **C0** corrected (profiles differ, TIS/MIS fixes it) / **C1** kernel-pinned (RL default; drift gated in CI) / **C2** bitwise (batch-invariant kernels + Behavior Record, for goldens and MoE parity). One repo ≠ automatic parity — the ladder is what makes the single-runtime bet honest.
5. **Extensions, never monkeypatching**: read-only observers (ParityAligner, ActivationTrace, Profiler, NaNWatch) and compute-altering interceptors at declared points — **cache-dit** is the first interceptor. **Dynamo is the fleet layer** (first-class partner, not a dependency we rebuild): registration/health/cost contract in-engine, seven concrete upstream asks (affinity key spaces, cost interface, media streaming, role-graph disagg, RL weight plane, KVBM generalization, sessions) — each with a fallback.
## Why this is the moat
Unlike LLMs — where inference optimization is post-hoc on frozen weights — **a usable video model is itself a post-training artifact**. Every inference capability we ship is a *(recipe, runtime)* pair: step distillation ↔ few-step samplers; self-forcing ↔ causal KV streaming; QAT-NVFP4 ↔ FP4 kernels; VSA ↔ sparse attention backend; RL ↔ samplers + capture. The training loop *embeds* the inference loop, so whoever owns both sides of the pair owns the optimization frontier. The industry's RL pain proves the converse: verl-omni re-implements Wan2.2 inside vLLM-Omni and corrects the numerics afterward; miles' headline features are all mismatch patches for two runtimes with different kernels. We answer with one model definition, one kernel set, one measured ladder — at FastVideo's 1–30B FSDP2 scale, where the bet is viable.
## What's pulling on it
| Customer | Pull | Proof point |
|---|---|---|
| **Cosmos3 / omni** | MoT loops, packed sequences, reasoner KV, world-model rollout | 150-test parity suite on `feat/cosmos3-reasoning` |
| **Dreamverse** | Engine-client replaces hand-rolled pool; capacity = duty cycle + cost-model admission + distillation | today 1 B200/session; Phase-3 gate: ≥2 sessions/GPU on a recorded duty-cycle trace, p95 within SLO |
| **RL (landed)** | #1450 migrates onto shared loops (Phase 1), engine-client rollouts (Phase 2+); GRPO-class next | the vendored-sampler docstring; C1-by-construction discipline |
| **ComfyUI funnel** | embed (nodes) → **compile** (workflow→PipelineSpec, accelerated cloud) → productize (Studio) | tier-1 ~20-node static sublanguage maps onto PipelineSpec |
## Migration — seven phases, each independently shippable
| Phase | Ships | Gate |
|---|---|---|
| **−1** | Merge cosmos3 chain; seed SSIM for uncovered families | baselines exist |
| **0** | Typed omni I/O; config freeze (`compat.py` shrinks monotonically to zero) | all SSIM suites unchanged |
| **1** | **Loop inversion** + policies + extension core (cache-dit, ParityAligner); RL migrates off its vendored sampler | old vs new loop **bit-identical**; per-method grad-norm refs (#1396) extended |
| **2** | AsyncEngine + StepScheduler; LTX-2 linear graph; Dynamo worker (stock); colocated weight sync | ≤2% batch-1 latency regression; Dreamverse single-session parity; RL engine-client parity |
| **3** | PipelineSpec graphs, role pools, declarative parallelism, ComfyUI compiler MVP, general WeightSyncPlan | ≥2 Dreamverse sessions/GPU on recorded duty-cycle trace |
| **4** | Cosmos3 native; AR continuous batching + paged KV (arriving *with* their workload, per N5); RL hardening (C1/C2, Behavior Record) | Cosmos3 parity suite on new runtime; drift ≈ 0 on a Wan RL run |
| **5** | Deletion: legacy `training/` retires, then `ComposedPipelineBase`, legacy loop, `forward_context.py`, `compat.py`, `RayDistributedExecutor` | the deletion diff — **4 loop copies → 1** |
Enforcement the last freeze lacked (it was broken 19×): `compat.py` frozen from Phase 0; CI path gates + CODEOWNERS once Phase 1 lands; new families land on new abstractions from Phase-1 completion.
## What we are deliberately not doing
Datacenter orchestration (Dynamo's job) · trainer internals (frozen `training/`; `train/` is a consumer) · replacing the bit-exact porting methodology · migrating 20+ families at once · **standalone LLM-serving excellence** — AR machinery arrives only at the sophistication omni workloads pull (N5).
## Decisions we need from this review
1. **sglang `multimodal_gen` relationship** — upstream, friendly fork, or shared core (decide by Phase 2; drift is a strategic cost either way).
2. **Dynamo asks** — green-light proposing the Phase-2 asks (A2 cost interface, A3 media streaming) to the team first, with A5 (RL weight plane) queued behind them?
3. **Phase −1 start** — merge the cosmos3 chain and seed SSIM baselines now; it blocks everything else.
-830
View File
@@ -1,830 +0,0 @@
# FastVideo v3 — A Model-Native Runtime for the (Recipe, Runtime) Era
**Status:** unconstrained north-star. This document assumes we are free to build a brand-new architecture with no
backward-compatibility, no migration tax, and no obligation to the current code. It exists to define the *ceiling*:
the system FastVideo should be if nothing held it back. Migration is a separate, later question — deliberately out of
scope here.
**Lineage:** this is the synthesis of `design.md` (the strategic thesis, product pull, and hard-won serving realism)
and `designv2.md` (the model-native center and typed contracts), with the open tensions of both resolved rather than
hedged.
---
## Table of contents
1. [The thesis](#1-the-thesis)
2. [The two signature ideas](#2-the-two-signature-ideas)
3. [Planes and their dependency order](#3-planes-and-their-dependency-order)
4. [Model Plane — the center](#4-model-plane--the-center)
5. [The loop contract — driven loops](#5-the-loop-contract--driven-loops)
6. [Runtime and scheduler — one currency, one WorkUnit](#6-runtime-and-scheduler)
7. [Memory, cache, transport, compile](#7-memory-cache-transport-compile)
8. [Parallelism as a model contract](#8-parallelism)
9. [Correctness — parity as a typed gate](#9-correctness)
10. [Training and RL on the same loops](#10-training-and-rl)
11. [Extensions — observers and interceptors](#11-extensions)
12. [Request, session, artifact, stream](#12-request-session-artifact-stream)
13. [Programs and workflows](#13-programs-and-workflows)
14. [Deployment and fleet](#14-deployment-and-fleet)
15. [Worked examples](#15-worked-examples)
16. [What this unlocks](#16-what-this-unlocks)
17. [Honest unknowns and falsifiers](#17-honest-unknowns-and-falsifiers)
18. [Package layout](#18-package-layout)
19. [Reference synthesis](#19-reference-synthesis)
---
## 1. The thesis
Three facts about video generation, taken together, dictate the architecture.
**A deployable video model is a post-training artifact.** Unlike an LLM — where inference optimizes frozen weights
post-hoc — a *usable* video model is *created* by training: step distillation is mandatory for usable latency, low
precision needs QAT, and causal/world models are made by distillation plus self-forcing. Every inference capability is
therefore a **(recipe, runtime) pair**: the weights and the loop that produced-and-assumes them are inseparable. A
"4-step NVFP4 FastWan" is not weights plus a flag; it is a distillation recipe, a sampler, a precision path, and a
parity contract that are one object.
**Video systems are loop systems.** Denoise timesteps, AR decode, chunked world-model rollout, VAE tiles, encoder
chunks, audio tokens, reward batches, optimizer steps, media chunks — the work is iteration, not a single `forward`. A
runtime that reduces everything to `forward()` cannot schedule, batch, cancel, stream, reserve memory for, or capture
the behavior of the thing that actually runs.
**Omni models share weights across loop types within one request.** Cosmos3's text reasoner and multimodal denoiser
are the *same resident weights*, driven by an AR loop and then a diffusion loop in one request. This cannot be a
DAG of separate engines (that doubles 30B+ of weights and severs the shared KV/denoise state); it must be one resident
model instance running many loop types. This is achievable — vllm-omni's `bagel_single_stage`/`lance` already run one
resident MoT instance doing AR `generate_text` and diffusion `generate_image` on co-resident experts in a single
request — *but* they bury that interleaving inside one opaque `DIFFUSION` stage their scheduler never sees inside,
request-scheduled with `max_num_running_reqs` forced to 1. The hard part, and the differentiation, is not *expressing*
the shared-weight loops — it is making them **runtime-visible, step-scheduled, batchable, and cost-priced.**
The architecture that falls out:
> **The atomic unit is the (recipe, runtime) pair, owned by a typed `ModelCard`.** Everything — serving, training, RL,
> products, deployment — is a *view* over that card. The **runtime owns loop *lifecycle*** (admission, scheduling,
> batching, caching, cancellation, streaming, behavior capture); the **model owns loop *semantics*** (typed state
> transitions and kernel execution). One resident model instance runs many loops; one scheduler schedules the
> *steps* of all of them in a single currency; one parity contract binds the train-forward to the serve-forward so
> the recipe and the runtime never silently drift apart.
The single invariant, stated once:
```text
Model cards own components, loops, recipes, and parity.
Programs compose loops into tasks.
The scheduler executes the steps of loops as WorkUnits under one budget.
Caches are correct by key, not by hope.
Training records behavior on the same loops it serves.
Deployment places and routes; it does not define semantics.
Products stream artifacts; they do not reach into the model.
```
Everything below is the elaboration of that invariant.
---
## 2. The two signature ideas
Two ideas do most of the work and are what an unconstrained design can reach that an incremental one cannot.
### 2.1 The (recipe, runtime) pair is a first-class, versioned, typed object
A model in v3 is not a checkpoint. It is a `ModelCard` that owns, as one versioned unit:
- the **components** (weights, loaders, layouts),
- the **loops** it can run (the runtime semantics),
- the **recipe** that produced the weights (distillation/QAT/RL config, teacher, data contract, the sampler the
recipe assumes), and
- the **parity contract** asserting that the train-forward and the serve-forward agree to a declared level.
You cannot ship the weights without the loop they assume, and you cannot change the loop without re-proving parity.
This makes design.md's "(recipe, runtime) pair" *literal*: the deployable artifact carries its own provenance and its
own correctness obligation. It is the thing that turns "we do training and inference in one repo" from an org chart
into a guarantee — and it is essentially un-retrofittable, which is exactly why it belongs in a clean design.
### 2.2 Driven loops: the model owns control flow, the runtime owns execution
Loop inversion, done right, is not a heavy `plan_step`/`run_step` contract and not a hidden `for t in timesteps`. It
is a **driven loop**: the model describes the *next step it needs*, the runtime *decides when and with whom that step
runs*, and the model folds the result back into its own state and decides what to do next. The model keeps its control
flow (so content-adaptive decisions — cache-dit skips, EOS, VSA tile selection — are natural); the runtime keeps the
`await` (so admission, batching, cancellation, streaming, and behavior capture are universal). Per-request state lives
in the loop's own typed `LoopState`, never in module globals, so interleaving requests through one model instance
cannot smear state — the failure mode that makes naive loop-inversion dangerous is *structurally* excluded.
These two ideas are developed in §4–§5. The rest of the system is their consequence.
---
## 3. Planes and their dependency order
```text
Products: Python · CLI · OpenAI API · ComfyUI · Dreamverse · RTC · Trainer
│ (thin: validate intent, make requests/sessions, subscribe)
Request / Session / Artifact / Stream
│ (typed runtime objects, cancellation, streaming)
Program Plane ← typed loop programs + compiled workflows
│
┌──────────────── Model Plane (CENTER) ────────────────┐
│ ModelCard: components · loops · recipe · parity │
│ capabilities · caches · parallelism · precision │
└───────────────────────┬──────────────────────────────┘
│
┌─────────────────────────────┼─────────────────────────────┐
│ Runtime / Scheduler │ Training / RL │ (same loops, different capture)
│ WorkUnits · GPU-time budget │ rollout · reward · weight-sync│
└─────────────────────────────┼─────────────────────────────┘
│
Memory · Cache · Transport · Compile (typed CacheKey, per-class pools, CuMem sleep/wake)
│
Parallelism (named axes → DeviceMesh, validated, part of the cache key)
│
Deployment / Fleet (DeploymentCard → Dynamo; never the core)
```
Dependency rules (enforced at the package boundary, §18):
- Products do not define model semantics. Workflows do not define model semantics. Deployment does not define model
semantics. Training does not redefine model semantics. **All of them reference the Model Plane.**
- The runtime *executes* model loops but does not *own* their math. Training *captures* behavior on serving loops but
does not *fork* them. Cross-cutting concerns — extensions (§11), parallelism (§8), and parity (§9) — are contracts
declared on the card, not features bolted onto the runtime.
---
## 4. Model Plane — the center
### 4.1 ModelCard
```python
class ModelCard:
model_id: str # "fastwan-1.3b-nvfp4-4step"
family: str # "wan"
components: dict[str, ComponentSpec]
loops: dict[str, LoopSpec]
capabilities: CapabilityMatrix # text_to_video, image_to_video, reasoning_text, vae_decode, ...
recipe: RecipeSpec # ← what produced these weights (signature idea §2.1)
parity: ParitySpec # ← train-forward ≡ serve-forward, to a declared level (§9)
caches: dict[str, CacheContract]
parallelism: ParallelismContract
precision: PrecisionContract
checkpoint: CheckpointManifest # explicit components, layouts, key maps — no name-detector guessing
```
The card is both a **declarative contract** (strict enough to validate before any GPU touches it) and a **runtime
factory** (it knows how to instantiate components, bind loops, and resolve caches). It is hub-interchange compatible
with diffusers' `modular_model_index.json` / `ComponentSpec` so models published either way load both ways.
`CheckpointManifest` replaces today's implicit `model_index.json` + name-detector resolution with explicit declared
components and `required_for` / `optional_for` task sets (the Cosmos3 lazy-sound-VAE problem becomes a declaration, not
an `if env_var` inside `forward`).
### 4.2 RecipeSpec — the provenance half of the pair
```python
class RecipeSpec:
method: str # "dmd2" | "self_forcing" | "attn_qat_nvfp4" | "diffusion_nft" | "base"
parents: list[str] # teacher / base model_ids this was distilled or RL'd from
data_contract: DataRef # what the recipe trained on (for governance and reproduction)
assumes_loop: str # the loop_id this recipe's weights require at serve time
assumes_precision: str # the precision the QAT recipe baked in
consistency_required: str # the minimum parity level this recipe's outputs are valid under (§9)
```
`assumes_loop` and `assumes_precision` are the teeth: a 4-step distilled model whose `assumes_loop = "ddim_4step"`
cannot be served under a 50-step sampler without a typed mismatch error. The recipe and the runtime are bound.
### 4.3 ComponentSpec and LoopSpec
```python
class ComponentSpec:
component_id: str
kind: str # dit | vae | text_encoder | reasoner_tower | reward_head | ...
load_id: str
config_schema: type
io_schema: tuple[type, type]
precision_policy: PrecisionPolicy
placement_policy: PlacementPolicy
parallel_constraints: ParallelConstraint
parity_tests: list[ParityTestSpec]
class LoopSpec:
loop_id: str # diffusion_denoise | ar_decode | chunk_rollout | vae_tile | ...
state_schema: type # the typed LoopState (no dicts)
step_schema: type # the typed WorkPlan a step emits
result_schema: type # the typed StepResult a step returns
behavior_schema: type | None # what to capture for RL (None if not training-relevant)
step_cost_model: CostModel # predicted GPU-time per step at (shape, precision, policy) — §6
valid_parallel_plans: list[ParallelPlanPattern]
graph_capture: GraphCapturePolicy
cache_policy: CachePolicy
```
A `ModelInstance` is a resident, loaded card: component instances, model state, caches, compiled graphs, and a
parallel plan. **A request may run several of the card's loops against one `ModelInstance`.** That single sentence is
the difference between this design and a stage-only design, and it is what makes omni native:
```text
one Cosmos3 ModelInstance, one request:
ar_decode(reasoner) → pack → diffusion_denoise(vision[+action][+sound]) → vae_tile_decode → audio_decode
└────────────── same resident weights, shared packed state, scheduled as steps ──────────────┘
```
---
## 5. The loop contract — driven loops
### 5.1 The contract
A loop is a **serializable state machine** the runtime drives:
```python
class Loop(Protocol):
def init(self, req: Request, model: ModelState, ctx: LoopContext) -> LoopState: ...
def next(self, state: LoopState) -> WorkPlan | Done: ... # describe the next step; NO GPU kernels here
def advance(self, state: LoopState, result: StepResult) -> LoopState: ... # fold result in; decide what's next
def finalize(self, state: LoopState) -> LoopResult: ...
```
The runtime's driver — the only place iteration lives:
```python
state = loop.init(req, model_state, ctx)
while True:
plan = loop.next(state) # typed WorkPlan: resources, cache reads/writes, shape, sinks, cancel-scope
if isinstance(plan, Done):
break
result = await ctx.execute(plan) # ← THE INVERSION POINT: runtime admits, batches, places, runs, returns
state = loop.advance(state, result) # content-adaptive: next() can branch on everything in state, incl. result
for chunk in plan.emits:
ctx.emit(chunk) # streaming falls out
return loop.finalize(state)
```
Why this is the right contract, point by point against the failure modes:
- **Content-adaptive steps are natural.** `next()` reads `state`, and `advance()` has already folded in the last
`StepResult` — so cache-dit's skip decision (a residual comparison from the prior step), AR's EOS, and VSA's
content-dependent tile selection are ordinary control flow in the model. This is the tension `designv2.md`'s
"pure plan_step that pre-declares shape" could not resolve; here it dissolves, because planning the *next* step is
allowed to depend on the *previous* result. `next()` is still kernel-free (it *describes* work; it does not run it),
which is all the scheduler needs.
- **Cross-request state safety is structural.** All per-request state is in the loop's typed `LoopState`. There are no
module-level residual/KV globals (the bug that silently corrupts cache-dit and TeaCache forks under concurrency).
Interleaving requests through one `ModelInstance` cannot smear state because there is no shared mutable state to
smear. This is the safety property naive loop inversion lacks, made impossible-to-get-wrong by construction.
- **The runtime owns everything it needs and nothing it doesn't.** `execute(plan)` is the single seam for admission,
memory reservation, cross-request batching, placement, graph dispatch, cancellation, and behavior capture. The model
never sees the scheduler; the scheduler never sees the model's math.
- **Serializable, therefore migratable and resumable.** `LoopState` is typed and serializable, so a half-finished
1000-step job is a resume point: preempt by stopping the driver, migrate by shipping `LoopState` to another worker,
recover from a crash by replaying from the last serialized state. (A coroutine that keeps state in a suspended Python
frame — the tempting sugar — cannot do this; the explicit state machine is the price of resumability, and it is
worth paying.)
Custom step bodies are first-class, not an escape hatch with an asterisk: a family whose math is genuinely braided
(Cosmos's EDM coefficients consumed inside the CFG branch with an x0-space combine; LTX-2's 1–4 runtime-decided
guidance passes) writes `next()`/`advance()` by hand using samplers and CFG utilities as a *library*. The runtime
requires only the four methods; *how* a step body is factored is the model's business. Policies (CFG, expert routing,
precision, flow-shift, conditioning) are the *default* decomposition that deletes duplication for the families that
fit — never an admission requirement.
### 5.2 Loop granularity
Chosen by runtime value, not purity:
- too coarse → cannot cancel/batch/reserve/stream/record at useful points;
- too fine → scheduler overhead dominates, graph capture fragments;
- good default → one denoise step (or window), one AR decode batch, one encoder chunk, one VAE-tile batch, one
reward/logprob batch.
The runtime may **fuse** adjacent compatible WorkPlans after planning (an optimization); the unfused boundary remains
the semantic model, so parity and behavior capture are defined on the unfused loop.
### 5.3 CFG is a policy over *one* shared denoise body (verified)
A natural worry: CFG changes the *shape* of the step (one forward vs two vs a batched pair vs a data-parallel split),
so can one shared denoise loop really host all of it by swapping a policy? **Yes — and it is proven by existing code,
not aspiration.** vllm-omni's `CFGParallelMixin.predict_noise_maybe_with_cfg` + `combine_cfg_noise`
(`diffusion/.../cfg_parallel.py:76-212`) already runs sequential-2-forward, batched-1-forward, *and* cfg-parallel
through **one** pair, with the loop body unaware of which. The clean cut is three layers:
- **In-loop `CFGPolicy`** — branch vocabulary (`[cond]`, `[cond, uncond]`, per-modality, STG-perturbed),
the combine formula (standard `uncond + s·(cond−uncond)`, CFG-zero `st_star`, `cfg_normalize`/`guidance_rescale`),
and **per-request mutable state** (the adaptive-gate cached delta with model-id self-invalidation is the canonical
state case — and exactly why state lives in `LoopState`, §5.1). **Batched-vs-two-forward is a *dispatch detail
inside one policy*, not a separate mechanism.** This covers classic / batched / adaptive-gate / per-modality.
- **`cfg`-parallel is a *parallelism axis*, not a policy** — it shards the policy's branches across ranks and runs the
*same rank-invariant `combine`* on every rank. It composes *under* any `CFGPolicy`; you own a `BatchedCFG` policy
**or** a `cfg` group, never both (the §9 build-guard).
- **Companions are an *orchestrator pattern*, not in the loop** — splitting a request into companion sub-requests
upstream of diffusion (the conditioning is precomputed and bundled in; the loop is unchanged).
Two caveats keep the first pass honest: the `combine` runs in *the step body's* numeric space (Cosmos combines in
x0-space after EDM preconditioning, not noise-space — the body fixes the space, the policy fixes the algebra), and
embedded-guidance (Flux) is a **degenerate single-branch identity-combine policy** (guidance rides inside the forward
kwarg), kept *inside* the same abstraction rather than special-cased as "no CFG." This is the same shared denoise body
that RL rollout reuses (§10) — one CFG taxonomy serves both serving and rollout.
---
## 6. Runtime and scheduler
### 6.1 One WorkUnit, one currency
Every `await ctx.execute(plan)` produces a **WorkUnit**: the smallest schedulable action with a resource reservation
and a loop boundary. Kinds: `ar_prefill`, `ar_token`, `diffusion_step`, `diffusion_window`, `chunk_step`,
`encoder_chunk`, `vae_tile`, `audio_chunk`, `reward_batch`, `logprob_batch`, `transfer`, `cache_io`, `graph_capture`.
Tokens are *one kind*, not the scheduler — this is the generalization of vLLM's token scheduler that diffusion forces.
**The budget currency is predicted GPU-time, not counts.** A bidirectional denoise step re-attends the full latent at
O(L²) with zero KV amortization (every step pays full price); an AR decode step is ~O(context) against a cache; a
chunked-causal step sits between. Counting "steps" or "tokens" puts items three orders of magnitude apart in one
bucket. So each WorkUnit converts to GPU-seconds via its `LoopSpec.step_cost_model`, calibrated online by the Profiler
(§11). **The same cost model is the interface published to the fleet** (§14): the scheduler's internal budget and
Dynamo's routing/autoscaling input are one object, built once.
Two honesty caveats, kept from design.md's contact with reality:
- **Admission uses the conservative baseline.** The design's own flagship features make realized cost unknowable in
advance — cache-dit skips are residual comparisons, VSA tiles are content-dependent, AR length is unbounded
(budgeted at the `max_tokens` cap, refunded on early EOS). Telemetry refines calibration; it never licenses
admission optimism.
- **A denoise step is indivisible.** A 30s-1080p step bounds iteration latency no matter the budget. Mitigations are
first-class, not afterthoughts: cost-class pools (jumbo steps don't co-schedule with latency-class work), SP within
a node to shrink jumbo wall-time, and admission-time SLO classes so the fleet planner scales pools per class. This
is why the scheduler is *cost-aware*, not just *count-aware*.
### 6.2 WorkPlan and admission
```python
class WorkPlan:
loop_id: str
instance_id: str
kind: str
shape_sig: ShapeSignature # for batch compatibility + graph capture key
resources: ResourceRequest # compute (GPU-s), resident bytes, peak-activation bytes, cache blocks, xfer bw, sinks
cache: CachePlan # typed reads/writes (§7)
placement: PlacementHint
cancel_scope: CancelScope
emits: list[StreamChunk]
class Done: result: LoopResult
```
**Admission rule (the soundness condition of multiplexing):** *do not admit a waiting WorkUnit unless every resource
it requests can be reserved* — compute budget **and** memory (resident + worst-case peak) **and** cache blocks **and**
transfer bandwidth **and** graph-capture shape **and** output sinks. Two requests that fit individually but jointly OOM
are rejected at admission, not discovered at step 37. This is vLLM's "token budget is half the story, `allocate_slots`
is the other half," generalized.
### 6.3 The scheduler, in layers (each testable on a fake pool, no GPU)
1. **RequestScheduler** — accepts requests/sessions, selects programs, starts loop drivers.
2. **LoopScheduler** — drives `next()`, collects pending WorkPlans.
3. **BatchScheduler** — groups compatible WorkPlans by `(instance, loop_kind, shape_sig, precision, parallel_plan,
graph_key)`; image diffusion and AR decode batch across requests, jumbo video stays batch-of-1, Cosmos-style
token-budget packing is an opt-in.
4. **PlacementScheduler** — worker, role pool, instance, device mesh.
5. **TransferScheduler** — tensor / cache / artifact movement as scheduled WorkUnits.
6. **AdmissionController** — the reservation gate of §6.2.
Policies: running loops first (vLLM); preempt only at loop step boundaries; cancel only at declared scopes; prefer
cache hits when latency/fairness allow; never starve a long denoise loop behind short AR requests.
### 6.4 SPMD consistency and failure isolation
All ranks of a pool must make identical scheduling decisions or NCCL deadlocks. **Rank-0 decides, broadcasts** — the
existing discipline, now also the channel for the **abort broadcast** (failure isolation and scheduling share one
consistency mechanism). Failure classes: *request-fatal* (NaN flagged by NaNWatch, one request's step error) →
SPMD-consistent abort of that request, deliver partial artifacts with a structured error; *pool-fatal* (illegal
access, NCCL desync) → pool re-init, invalidate pool caches, resume requests from serialized `LoopState` where one
exists. **Cancellation is common-path, not exceptional** — vibe-directing makes abandoning in-flight work the *normal*
user action; cancel takes effect at the next step boundary, drops queued WorkUnits, releases `LoopState` and cache
handles, reports `cancelled`.
---
## 7. Memory, cache, transport, compile
Video and omni inference are memory systems as much as compute systems. This plane is explicit and typed.
### 7.1 Cache correctness is a contract
```python
class CacheKey:
model_id: str; component_id: str; loop_id: str | None
weights_version: str; adapter_versions: dict[str, str]
precision: str; parallel_plan_hash: str
shape_sig: str; layout_sig: str
scheduler_sig: str | None; guidance_sig: str | None; seed: int | None
input_hashes: dict[str, str]; step_index: int | None
contract_version: str
```
**If a field can change output semantics, it is in the key.** Incorrect reuse is worse than no reuse. The serving
hazard this kills: a workflow-cloud request that shares a prompt but differs in te-LoRA stack must not serve stale
embeddings — so the key is *partitioned* by `adapter_versions`, not flushed. An RL `update_weights` bumps
`weights_version` and invalidates wholesale.
### 7.2 Per-class pools (the granularity reality)
There is **no single unified block pool**, because a unified pool requires uniform bytes-per-block and our cache
classes differ by 150–500× in natural granularity (a text-KV page ≈ 64 KB/layer; a causal-video latent-chunk slab is
9.6–32 MB/layer) and their demand is workload-decoupled. Each class gets a statically budgeted pool behind one
`CacheHandle`: paged text-KV (`ar_decode`), slab chunk-KV (`chunk_rollout`, with a declared training mode that
disables mid-rollout recycling and keeps grad-aware index snapshots), feature caches (text/vision-encoder, content-hash
keyed, reference-counted FIFO), residual caches (cache-dit, scoped per `LoopState`), weight/adapter cache
(disk→CPU→GPU LRU for the workflow cloud). MoT falls out: the und pathway draws paged KV, the gen pathway draws slab or
nothing — independent budgets, no interference. **KV is the minority case** — a pure bidirectional deployment allocates
none of it; the machinery materializes only when a card declares KV-bearing loops.
### 7.3 Memory, transport, compile
- **Memory** — tagged pools, sleep/wake by tag (CuMem-style; tags are component names), reservation before admission,
per-role budgets, host-pinned staging. Sleep/wake is component-granular for RL (drop DiT + caches, keep
VAE/text-encoder resident).
- **Transport** — manifest-based and pluggable: in-proc reference → SHM → CUDA IPC → NCCL/UCXX/NIXL/RDMA →
object-store. KV/cache-bearing edges speak a `KVConnector`-shaped protocol (scheduler-side query/alloc/finish +
worker-side async load/save) so NIXL/LMCache/Mooncake/KVBM implement it directly. Transfers are scheduled WorkUnits,
not side effects.
- **Compile** — CUDA graphs and `torch.compile` managed by a `CompileCache` keyed on `(model, component, loop,
work_kind, shape_sig, precision, parallel_plan, backend)`. **Never full-graph across the engine** (per vLLM's own
reversal): per-block compile where it pays, manual fused ops permitted in model code, breakable CUDA graphs as an
*optimization tier* over an always-correct eager baseline. Graph capture is planned by the scheduler (padding,
bucketing, capture sizes affect admission and batching).
---
## 8. Parallelism as a model contract
Parallelism is not a launch flag; it affects cache keys, scheduling, transport, capture, and parity, so it lives on
the card.
```python
class ParallelPlan:
axes: dict[str, int] # dp, tp, sp(=ulysses×ring), cp, cfgp(≤2), pp_patch, vae, ep, fsdp, role, replica
mesh_order: list[str]
placement: PlacementSpec
communication: CommunicationSpec
```
Declarative, validated, compiled to a PyTorch `DeviceMesh` via a `ParallelDims`-style builder
(product-of-degrees validation, cached submeshes). **Pre-flight or it fails at load, never halfway.** Ownership
conflicts are build errors (CFG owned by a `BatchedCFG` *policy* or a `cfgp` *group*, never both). Applicability
conditions travel with axes: `pp_patch` (PipeFusion displaced-patch pipelining) is **invalid for causal/AR** (stale KV
breaks causality) and the validator enforces it per card. Degree-one axes exist as trivial groups so component code
needs no special cases. **Pools are single-node**; multi-node scale is *multiple pools* fronted by the fleet (§14) —
the engine never owns a cross-node NCCL mesh inside one pool.
---
## 9. Correctness — parity as a typed gate
This is the section both prior documents needed and neither fully had.
### 9.1 The parity contract
Every card carries a `ParitySpec`. Parity is **measured, never assumed**, by a `ParityAligner` observer (§11): record
named taps per step/block from a reference (the official framework, or a pre-change build); compare-mode replays with
fixed seeds and reports the first divergence beyond per-tap tolerance. This is the engine behind the "old loop vs new
loop, bit-identical" gate and the standing instrument for every port, precision change, and kernel swap.
### 9.2 The consistency ladder (with the rung both prior docs missed)
```text
C0 component parity — VAE, encoder, transformer block, scheduler step in isolation
C1 loop parity — full denoise trajectory / AR logits, fixed seed
C2 behavioral identity — the train-forward and serve-forward agree on the quantity the RL objective uses:
· likelihood-based methods (GRPO-class): per-step log-prob identity
· likelihood-free methods (DiffusionNFT-class): seeded final-sample +
prediction-space identity (old_deviate / ref-MSE) — there are NO log-probs to match
C3 distribution parity — rollout distribution under allowed nondeterminism
C4 artifact quality — SSIM-class, reward agreement, human-preference (gates product claims; needs the eval system)
```
The C2 split is load-bearing and is the lesson of the landed RL stack: the shipped Wan DiffusionNFT is
**likelihood-free** — it captures only final clean latents and contrasts the student against an implicit negative
policy in prediction space, so "log-prob identity" is *undefined* for it. A ladder that assumes log-probs (as both
`design.md`'s and `designv2.md`'s early framings did) cannot describe the only RL method actually in the tree. RL
methods declare their required level on the `RecipeSpec`.
### 9.3 The gate that catches what batch-of-1 cannot
Loop inversion's real hazard is **cross-request state smearing under interleaving** — and a batch-of-1 parity gate is
*structurally blind* to it, because the corruption only manifests when two requests share a loop. §5.1 excludes the
hazard by construction (state in `LoopState`, never globals), but construction-arguments need a test. So v3 makes a
**batch-of-N interleave parity test** a *required* gate: two (or more) concurrent requests, interleaved at step
granularity, must be bit-identical to the same requests run serially. This is the test the whole loop-inversion bet
lives or dies on, and it is named here as a first-class obligation, not left implicit.
### 9.4 Three execution profiles, one definition
Even in one runtime there are three forwards: the **serve** forward (no-grad, graphed, cached, possibly quantized), the
**rollout** forward (serve profile + behavior capture), and the **train** forward (grad, checkpointed, FSDP-gathered).
They share *one* loop definition; they differ only in grad mode and capture. The ladder measures the gap; the recipe
declares the level it needs. "Train BF16, serve FP8" is legal only at C2-corrected with importance-sampling, and the
card says so. This is how the (recipe, runtime) pair stays honest: the contract is typed and tested, not trusted.
---
## 10. Training and RL on the same loops
```text
serve : request → program → loop → WorkUnits → artifacts
rollout : prompt batch → program → loop → WorkUnits → BehaviorRecords → rewards → update
```
The loop kernel is shared; the only difference is output capture and training policy — not a second interpretation of
the model. This is design.md's §8 thesis and v2's training plane, with the dependency rule kept absolute: **`training`
may require behavior records but must not fork serving loop logic; the engine never imports `training`.** The engine
*is* the rollout engine (it already runs the loops); the trainer is a client.
**This is the moat — and it is the one place a serving-only runtime structurally cannot follow.** vllm-omni proves
omni serving can be production-grade, but it has *no* training/RL plane at all; verl-omni and miles prove the
alternative — a standalone trainer-side sampler on a *different* runtime than serving — costs the two-runtime tax
forever. The whole point of collocation is that **the rollout forward *is* the serve forward plus capture**: same loop,
same caches, same batcher, same numerics. Three consequences nothing else gets:
- **Every serving optimization is automatically a rollout optimization.** Distilled few-step samplers, cache-dit
skips, CFG-parallel, paged/feature caches, step batching — the recipe team builds them once for serving and the RL
rollout inherits them for free. FastVideo's *own* landed DiffusionNFT is the negative example that proves the
point: it vendors a bare-model `for`-loop (`rl/common/sampling.py`, whose docstring says it "intentionally does not
call FastVideo's full inference pipelines"), and DMD2 vendors a *second* one (`dmd2.py::_student_rollout`) — so
today's rollout runs with **zero serving-grade optimizations** (no CFG, dense attention, full 25-step ODE,
one-sample-at-a-time). Collocation deletes both private loops.
- **RL rollout is a *better* batching case than open-world serving — not a worse one.** A GRPO/NFT group is K
*identical-config* samples of one prompt: same shape, same schedule, same CFG branch. The landed config is K=24
(`num_video_per_prompt: 24`), 6 prompts/batch × 48 batches = **288 prompt-slots/GPU/epoch, each a 24-wide homogeneous
denoise batch** — zero bucketing required (serving must bucket heterogeneous resolutions/steps/CFG across users; a
GRPO group is homogeneous *by construction*). And all K samples share one prompt embedding, so the content-hash
feature cache computes the text encoder **once per group instead of 24×**. The vendored sampler captures none of
this; it carries the embedding per sample and runs one shape at a time.
- **One numerics surface.** Serve-forward and rollout-forward differ only in grad mode and capture (§9.4), so there is
no rollout-vs-train kernel gap to patch — the consistency ladder *measures* the gap rather than a correction layer
*papering over* it. For the landed likelihood-free NFT, "reuse holds" means it holds at the **C2 behavioral rung**
(seeded sample + prediction-space identity), under a `CFGPolicy` that is conditional-only and a `WeightSyncPlan`
whose role is the decay-blended old policy — all of which the card already declares.
- **BehaviorRecord** — captured at generation time (reconstructing later is fragile): seeds, scheduler trajectory,
timesteps, latents-or-refs, logprobs *where applicable*, sampled/action tokens, guidance, reward in/out, cache
assumptions, precision, parallel plan, attention backend, deterministic flags, `weights_version`. Sized honestly:
full MoE-routing capture is GB/sample for Cosmos3-class requests, so it is an **opt-in instrument** for goldens and
debugging, not always-on.
- **Weight-sync lifecycle** — freeze admission for the affected role/version → drain or boundary-stop in-flight loops →
transfer weights/deltas → bump `weights_version` → invalidate incompatible caches and graphs → publish version →
resume. A `WeightSyncPlan` is three inputs (mesh specs + per-model layout adapters + transport), validated
pre-flight, CPU-testable on fake pools. RL ships a *role*, not "the weights": student / EMA / decay-blended old
policy is declared (the landed NFT behavior policy is the *old* copy, not the student — the plan must carry that).
- **Roles** (policy, rollout, reference, reward, critic, evaluator, data, coordinator) reference the same cards and
loops; they are deployment concerns, scaled by the fleet.
- **The industry tax we delete:** verl-omni re-implements Wan inside vLLM-Omni and corrects numerics afterward; miles'
headline features (TIS/MIS, bitwise logprobs, R3 routing replay, unified FP8) are all mismatch patches for *two
runtimes with different kernels*. One model definition, one kernel set, one measured ladder is the answer — viable
at FastVideo's 1–30B FSDP2 scale (the boundary condition: a Megatron-class trainer at 100B+ re-enters the
two-runtime world, and the ladder is the fallback there).
---
## 11. Extensions — observers and interceptors
The optimization, debugging, and parity surface, as versioned hook points assembled at loop build (an unused hook is
*literally absent* from the hot path). It composes with §5 cleanly: the hooks wrap `ctx.execute(plan)`.
- **Observers (read-only):** `ParityAligner` (§9), `Profiler` (per-step wall+CUDA, calibrates the cost model),
`NaNWatch` (first-NaN localization), `ActivationTrace`. They cannot mutate state.
- **Interceptors (compute-altering):** `StepInterceptor` (step-skip / cached-prediction) and `BlockInterceptor`
(cache-dit's DBCache/FBCache/TaylorSeer). State lives in `LoopState.plugin_state[id]`, keyed **per request and per
CFG branch** — the structural fix for the module-global residual state that silently corrupts cache-dit/TeaCache
forks under concurrency. cache-dit is the reference integration (the library sglang's serving already uses);
conflicting interceptors are rejected pre-flight; a 4-step distilled card *rejects* step-skip caches rather than
producing garbage.
**Trust boundary:** plugins are enabled at deploy scope only (never a per-request `plugins=[...]` field that would wire
third-party code selection into the public API); requests only *parameterize* pre-enabled plugins through validated
schemas, and exact-mode requests reject `distribution_altering` parameterization outright.
---
## 12. Request, session, artifact, stream
Typed runtime objects, not IDs in a batch (the Dreamverse/LiveKit lesson):
- `Request` — one generation, scoring, encoding, training-sample, or conversion job.
- `Session` — a long-lived interactive context: prompt memory, media streams, cancellation, partial updates,
cross-request chunk-KV that persists for a game/scene session.
- `Artifact` — a *named, typed* output with provenance (which node produced it): `VideoArtifact`, `AudioArtifact(
sample_rate)`, `TextArtifact(token_ids, text)`, `TensorArtifact`, `LatentArtifact`. This kills the `extra["audio"]`
pattern — audio carries its sample rate as a first-class artifact, not a dict passenger.
- `Stream` — one ordered event channel for previews, media chunks, progress, logs, finals.
- `CancelScope` — structured cancellation target (request / loop / stream / session).
Typed event taxonomy (`request.*`, `session.*`, `artifact.*`, `media.{init,chunk,complete}`, `trace.*`). A
`media.chunk` must know its stream, byte-range or shared-buffer ref, codec/container, timestamp range, and
preview-vs-final — invalid combinations are unrepresentable.
The **request is the only currency crossing the product boundary.** A typed `Request` carries `task: TaskType`
(declared, never inferred), `inputs: list[ModalPart]` (Text/Image/Video/Audio/Action/Latent), AR `sampling` vs
`diffusion` params, an `OutputSpec` (requested modalities + streaming + capture flags), and per-node overrides. Task is
declared; heuristics may only *suggest* a default at the boundary.
---
## 13. Programs and workflows
A **Program** composes a card's loops into a task; the card says what loops *exist*, the program says how to *run* them
for this request. Kinds: `InlineProgram` (many loops, one resident instance — the omni default), `DisaggregatedProgram`
(encoder→denoiser→decoder role pools), `WorkflowProgram` (compiled from ComfyUI), `TrainingProgram`,
`RealtimeProgram`. Nodes: `ModelLoopNode`, `ComponentNode`, `ExternalNode`, `ArtifactNode`, `ControlNode`,
`StreamNode`, `TransferNode`. Edges are typed (`TensorEdge`, `ArtifactEdge`, `StreamEdge`, `ControlEdge`, `CacheEdge`,
`BehaviorEdge`). Linear pipelines are the degenerate case; branches/fan-out/fan-in are real (video and audio decode in
parallel after a joint denoise). A separate deploy config maps nodes → pools/devices/parallelism, defaulting to "one
pool, everything colocated."
**Workflows compile, they are not the runtime.** A ComfyUI workflow's tier-1/tier-2 static sublanguage maps onto a
`Program` (`CheckpointLoaderSimple→card`, `KSampler→diffusion_denoise` with sampler/CFG policies,
`LoraLoader→adapter hot-swap`, `ControlNetApply→ConditioningInjector`); unknown nodes become `ExternalNode`s or a
coverage rejection — never silent wrongness. The moat: an orchestrator can run stock workflows on rented silicon;
substituting a *credibly faster* model requires owning the recipe (§2.1) — which an orchestrator structurally cannot
do. Equivalence is a quality-metric vs a reference render (C4), never a bit-parity claim.
---
## 14. Deployment and fleet
The engine exports a `DeploymentCard` and lets a fleet orchestrator (Dynamo) route — Dynamo orchestrates engines, it
is never the engine core.
```python
class DeploymentCard:
engine_id: str; model_cards: list[str]
capabilities: CapabilityMatrix; role_pools: list[RolePoolSpec]
supported_programs: list[str]; supported_parallel_plans: list[ParallelPlan]
cache_events: list[CacheEventSpec]; transfer_endpoints: list[TransferEndpoint]
cost_model: CostModel # the SAME §6 cost model — one object, two consumers
health: HealthSchema; slo: SLOSchema
```
Clean line: the **fleet** owns global routing, tenant policy, cold start, role-pool scaling, cross-node transfer,
placement-by-SLO, health/failover, multi-engine upgrades, global cache routing. The **engine** owns model load, loop
execution, local scheduling, local memory/cache, model-specific behavior, parity, WorkUnit batching. The asks of the
fleet are concrete and each has a fallback: generic affinity key-spaces (checkpoint/session/lora/weight_version beyond
token prefixes), a heterogeneous request-cost interface (the §6 cost model), chunked media streaming through the
frontend, role-graph disagg (N roles, not two), an RL weight plane (versioned broadcast + staleness-aware routing),
cache-object tiering (KVBM generalized to latent/session caches), and session lifecycle as a routing primitive.
---
## 15. Worked examples
**(a) Text → video, one instance.** `Request(T2V)` → `InlineProgram` → `diffusion_denoise` loop. Driver: `init`
builds sigmas/latents; `next` emits a `diffusion_step` WorkPlan (batch-of-1, SP+CFG-parallel); `advance` folds the
model output and the CFG combine; cache-dit's `BlockInterceptor` may skip blocks based on the prior residual; at
`Done`, a `vae_tile_decode` loop runs; output is a named `VideoArtifact`. Compiles/captures exactly like today's inner
loop — nothing tensor-level changes for batch-of-1.
**(b) Cosmos3 omni, one request, shared weights.** `InlineProgram` over one `ModelInstance`: `ar_decode(reasoner)`
yields `ar_token` WorkUnits that join the AR continuous-batching group → `pack` → `diffusion_denoise(vision+action+
sound)` yields `diffusion_step` WorkUnits → fan-out `vae_tile_decode` + `audio_decode`. The reasoner's tokens and the
denoiser's steps hit the *same resident weights*; the scheduler is the mode multiplexer. AR decode runs data-parallel
across the cfg×sp weight-replica axes (decode is sequence-length-1; SP has nothing to shard). This is the workload no
DAG-of-engines can express.
**(c) Image serving at scale.** Many `Request(T2I)` → the `BatchScheduler` groups `diffusion_step` WorkUnits by
resolution bucket and batches across requests every step — the case where cross-request batching pays most. The
*same* scheduler that runs (a) and (b).
**(d) RL rollout.** A `TrainingProgram` drives the *same* `diffusion_denoise` loop with `OutputSpec(capture=behavior)`;
each step emits a `BehaviorRecord` slice; rollouts run C2 by construction (in-process, trainer kernels, pinned
attention). For likelihood-free NFT the behavior is seeded final latents + prediction-space deviations; for a future
GRPO-class method it is per-step log-probs — the loop is identical, the capture differs, the ladder rung is declared.
**(e) Dreamverse session.** A `Session(realtime_video_continue)` holds chunk-KV across 5s segments; `push_text`
updates prompt memory; `stream` yields `media.chunk` previews from loop `emit`s; a direction change throws
`Cancelled` at the next step boundary and starts a new segment. Capacity comes from duty cycle + cost-model admission +
distillation — interleaving is fairness, not throughput.
**(f) ComfyUI compile.** `workflow.compile(json)` → a `WorkflowProgram` of `ModelLoopNode`/`ComponentNode` over a
weight-fleet-cached card, with stacked-LoRA patch/unpatch priced by the §6 weight-transition cost — same runtime, new
frontend.
---
## 16. What this unlocks (the unconstrained payoff)
Things no incremental design — and neither prior document — could actually claim:
- **True omni/MoT serving.** One resident model, many loop types, scheduled at step granularity, in one request. Not a
monolith bypassing the abstraction (the Cosmos3 port's necessary hack), not a DAG doubling weights — native.
- **Train ≡ serve by construction.** Because rollout and serve are the *same loop*, the (recipe, runtime) flywheel is
real and measured, not aspirational: Dreamverse's directing sessions emit preference data → the RL plane → faster
distilled cards → a better product, with the ladder guaranteeing the preferences collected under the serving profile
transfer into training.
- **Real-time interactive omni.** The driven-loop contract + sessions + WebRTC frame/PTS streaming + step-boundary
cancellation make the <100ms motion-to-photon interactive world-model loop expressible in the same runtime that
serves batch T2V.
- **One substrate, three personas.** A research vehicle (new ports land as cards), a product engine (Dreamverse, the
workflow cloud), and an RL rollout engine — without three codebases. The dependency rules keep them from fusing into
mud.
- **Correctness you can sign.** A deployable card is a *(recipe, runtime)* pair with a typed parity obligation; "this
fast model is equivalent" is a claim with a test behind it, which is the one thing an orchestrator-without-recipes
can never say.
---
## 17. Honest unknowns and falsifiers
An unconstrained design is not an unfalsifiable one. The bets, stated with the experiment that kills each:
- **The novelty is concentrated and real.** Runtime-owned diffusion iteration has a **narrow, opt-in precedent** —
vllm-omni's `SupportsStepExecution` (`prepare_encode/denoise_step/step_scheduler/post_decode`,
`diffusion/models/interface.py:44-67`) is exactly runtime-owned diffusion iteration at step granularity, and maps
almost 1:1 onto our `init/next/advance/finalize` — but it is Qwen-Image-only and off in every shipped deploy. What
is unprecedented is making it the **always-on universal contract** *and* a fully general WorkUnit scheduler over
heterogeneous units. The risk is not the loop contract (a state machine is well-understood, and now demonstrably
shippable); it is whether step-level cross-request scheduling *pays* for video. **Falsifier:** publish a load profile and targets from a real duty-cycle trace; if step-level scheduling does
not beat a request-level baseline (≥2 concurrent sessions/GPU, p95 within SLO), the scheduler degrades to
request-level dispatch and the loop contract keeps only its streaming/cancellation/behavior seams — which still
justify it. The contract is safe even if the scheduling bet loses; that is the design's insurance.
- **The general WorkUnit scheduler may be over-general.** Scheduling VAE tiles, transfers, and graph-captures through
the *same* admission machinery as denoise steps is elegant and unproven. **Falsifier:** if, after Phase 2, the
non-diffusion/non-AR WorkUnit kinds (tile, transfer, cache_io) gain nothing from unified scheduling over a simple
in-loop call, collapse them back to in-loop operations and keep WorkUnits for the step-bearing kinds only.
- **Cost-model admission is a modeling bet.** It converges toward cost-class pool routing once the indivisible-step
reality is respected — which is close to what request-level pooling + a fleet planner already do. The fine-grained
interleave win has a *narrow* window (many small concurrent jobs); it should be argued on that window, measured.
- **The clean-slate premise is the elephant.** This document deliberately ignores migration. The org that would build
it broke its own freeze 19 times and ships 20+ families, a live product, and a landed RL stack. A clean-slate
rebuild is the highest-risk path that exists for *this* org; the responsible realization is to build v3 as a
*parallel* engine around one forcing-function card (Cosmos3), prove it on the parity ladder, then migrate families
onto it behind an adapter while everything keeps shipping — i.e., reach this architecture incrementally. That plan is
out of scope here by request; it is non-optional in reality.
- **Quality is unmeasured.** C4 (artifact quality / human preference) and the eval system it needs do not exist yet,
and they gate every product claim ("fast mode is equivalent", RL reward validity, distillation comparisons). Named
as a required, currently-absent subsystem, not assumed.
---
## 18. Package layout
```text
fastvideo/
card/ specs, components, loops, recipes, parity, checkpoints, capabilities # the Model Plane
loop/ driver, loopstate, workplan, policies (cfg, expert, precision, flowshift, conditioning)
runtime/ engine, scheduler/{request,loop,batch,placement,transfer,admission}, workers, events
cache/ keys, classes/{paged_kv, slab_kv, feature, residual, weight_fleet}, policies
memory/ allocator, sleep_wake, reservations
transport/ manifests, backends/{shm, cuda_ipc, nccl, nixl, kvbm}, relay
parallel/ plans, mesh, process_groups, validation
parity/ aligner, ladder, interleave_gate # §9 is its own home
extend/ observers, interceptors, cache_dit, registry, trust
program/ specs, compiler, workflows
request/ requests, sessions, artifacts, streams, cancel
training/ rollout, behavior, rewards, weight_sync, methods # imports card/loop/runtime; never imported by them
deploy/ cards, role_pools, dynamo_adapter
integrations/ comfyui, dreamverse, livekit, diffusers
```
Enforced boundaries: `card/` imports no product/runtime; `runtime/` executes `card/` loops but defines no semantics;
`training/` may require behavior records but forks no loop; `integrations/` adapt external systems into core specs and
events, never bypass them. **`parity/` is a first-class package**, not a test folder — it is how the (recipe, runtime)
pair is kept honest.
---
## 19. Reference synthesis
| Source | Take | Constrain / reject |
|---|---|---|
| Cosmos3 (official + port) | Shared model instance across reasoning/diffusion/action/sound; packed multimodal sequences; component+scheduler parity matrices | A strong `ModelCard`, not the framework; no Cosmos-specific branching in global runtime |
| vLLM core | Running-first scheduling, reservation-before-admission, model-owned state, encoder/KV cache managers, CuMem sleep/wake, CUDA-graph dispatch, KV-connector split | Token scheduling is one WorkUnit kind; never full-graph compile |
| sglang `multimodal_gen` | Role pools, request lifecycle, capacity dispatch, transfer manifests, disagg state machine, cache-dit integration | No large mutable `Req`/`ForwardBatch` as the stable API; not single-item diffusion scheduling |
| vLLM-Omni | Frozen pipeline spec separate from deploy YAML (verified, adopt); `OmniConnectorBase` + `chunk_ready` readiness; **`SupportsStepExecution` as loop-inversion prior art** (opt-in, Qwen-Image-only — we generalize to always-on); TP-rank- and CFG-branch-aware KV-copy transfer; **`CFGParallelMixin` proves CFG-as-policy over one shared denoise body** (§5.3); 3 separate cache subsystems confirm per-class pools | Expresses shared-weight MoT (`bagel`/`lance`) only as **one opaque request-scheduled stage** the scheduler never sees inside — no step visibility, no cross-request batching by default; cross-stage KV is a *copy*, not a shared live cache; **no cost model** (per-stage count budgets); readiness-parking, not credit flow; RDMA = Mooncake/Mori/Yuanrong, not NIXL/NCCL |
| sglang-omni | The `next/wait_for/merge_fn/stream_to` edge vocabulary; Relay transport + **credit-based flow control** (this is sglang-omni's, not vllm-omni's) | Stages own disjoint weights; hybrid AR+diffusion only as AR-stage → DiT-stage; per-model bootstrap duplication |
| Dynamo | Fleet routing, disagg role pools, KV-aware routing, KVBM, SLA planner, ModelExpress cold-start/weight streaming | Orchestrates engines; never the engine core. Export a `DeploymentCard` + cost model to it |
| diffusers Modular | `ComponentSpec`/`modular_model_index.json` interchange; Guiders ≈ CFG policies | A Python pipeline interpreter is not the performance boundary; import is lossy |
| xDiT | DiT parallelism catalog (USP, ring/ulysses, PipeFusion, CFG-parallel, DistVAE) + world-size validation | Parallelism lives in the runtime + card, not a wrapper-per-model library; `pp_patch` invalid for causal |
| TorchTitan | Named mesh axes, `ParallelDims` validation, ModelSpec discipline, TorchStore weight-sync, batch-invariance utils | Adopt the discipline, not the stack; DCP/TorchStore don't reshard — `WeightSyncPlan` owns layout |
| verl-omni / miles / cosmos-rl | Rollout adapters, per-step capture, async rewards, group-relative advantage, TIS/MIS, deterministic/batch-invariant modes, per-payload `weight_version`, AIPO/off-policy masking | The two-runtime tax is the thing to delete; capture behavior *in* the serving loop, not after the fact |
| ComfyUI | Workflow graph, node-signature cache, model memory management, App-Mode (workflows-as-products) | Compile to `Program`; dynamic node execution is not the serving/training core; GPL hygiene |
| Dreamverse | Sessions, prompt memory, typed media IPC, cancellation, the duty-cycle capacity reality, the preference-data flywheel | Product/session behavior is first-class in the request plane, never merged into the model core |
| LiveKit | Realtime sessions, push audio/video, interruptions, turn/activity state, frame+PTS streaming | Realtime triggers only when they fire (<100ms interactive); don't force RTC onto offline jobs |
| Thinking Machines (batch-invariance) | The C2/C3 mechanism: batch-invariant kernels for bitwise rollout↔train identity | Scoped to goldens; the conservative baseline governs admission |
---
## Final position
```text
A model card is a (recipe, runtime) pair with a parity obligation.
The model owns loop semantics; the runtime owns loop lifecycle.
One resident instance runs many loops; one scheduler runs their steps in one currency.
Caches are correct by key; parity is correct by test; the interleave gate is non-negotiable.
Training records behavior on the same loops it serves.
Deployment places and routes; products stream artifacts; neither defines the model.
```
This is the ceiling: a model-native runtime where omni is native, train and serve are the same loops by construction,
correctness is a typed contract you can sign, and the (recipe, runtime) flywheel is real. The constraint we removed to
see it was migration. Putting that constraint back is the next document, not this one.
-1895
View File
File diff suppressed because it is too large Load Diff
+2 -4
View File
@@ -55,14 +55,12 @@ 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.
# 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).
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
./build.sh
EXPOSE 22
+2 -4
View File
@@ -55,14 +55,12 @@ 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.
# 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).
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
./build.sh
EXPOSE 22
+2 -4
View File
@@ -55,13 +55,11 @@ 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.
# 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).
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
./build.sh
EXPOSE 22
+2 -4
View File
@@ -55,14 +55,12 @@ 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.
# 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).
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
TORCH_CUDA_ARCH_LIST=9.0a ./build.sh
./build.sh
EXPOSE 22
@@ -108,7 +108,6 @@ 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."
@@ -126,18 +125,9 @@ 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
@@ -501,8 +491,6 @@ 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: {}
-4
View File
@@ -58,7 +58,6 @@ 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 | ❌ | ✅ | ✅ | ⭕ | ⭕ |
@@ -79,9 +78,6 @@ 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.
-124
View File
@@ -1,124 +0,0 @@
# 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()
@@ -1,98 +0,0 @@
# 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()
@@ -1,289 +0,0 @@
# 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()
@@ -1,350 +0,0 @@
# 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()
@@ -1,38 +0,0 @@
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()
@@ -1,91 +0,0 @@
"""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()
@@ -1,52 +0,0 @@
{
"data": [
{
"caption": "gold tip pyramid in the night, extremely detailed , rain, stars"
},
{
"caption": "a fruit stacking in the shape of a dog, stock image, shutterstock"
},
{
"caption": "Landscape, By Lee madgwick, by Luis Royo, by Louise nevelson"
},
{
"caption": "A colorful poster that says \"philo is a weird\""
},
{
"caption": "Danish male with blue eyes, realistic, viking"
},
{
"caption": "futuristic, cityscape, flying cars, neon lights, towering skyscrapers, glowing purple sky."
},
{
"caption": "a crow with cameras for eyes, sitting on a mans shoulder, anime, studio ghibli, fantasy, fairytale, sketch, digital art, watercolor, dnd, rustic, professional photograph, medieval, hd, 4k"
},
{
"caption": "a background image mixing the matrix and AI"
},
{
"caption": "Golden sunset, a bright orange and yellow sky is visible, lit up by the setting sun, the horizon is a mix of bright colors and deep shadows"
},
{
"caption": "Grim reaper playing an electric guitar"
},
{
"caption": "an epic view of a demonic Rose-ringed parakeet cyborg inside an ironmaiden robot,wearing a noble robe,large view,a surrealist painting, aralan bean and Philippe Druillet,hiromu arakawa,volumetric lighting,detailed shadows"
},
{
"caption": "Ben Shapiro as the cover of ministry's filth pig album, but covered in milk"
},
{
"caption": "king charles spaniel with , ethereal, midjourney style lighting and shadows, insanely detailed, 8k, photorealistic"
},
{
"caption": "A website for a party resort service"
},
{
"caption": "full shot of a steampunk horse"
},
{
"caption": "60s psycedelic spiritual jazz album art"
}
]
}
-1
View File
@@ -87,7 +87,6 @@ training:
# --- training.data [TYPED] -> DataConfig ---
data:
data_path: data/my_dataset # default: ""
preprocessed_data_type: t2v # default: "t2v" ("text_only" for simulate-only DMD text prompts)
train_batch_size: 1 # default: 1
dataloader_num_workers: 4 # default: 0
training_cfg_rate: 0.1 # default: 0.0
@@ -1,116 +0,0 @@
# DiffusionNFT multi-reward single-frame RL: Wan 2.1 T2V 1.3B on text-only PickScore prompts.
#
# Single-frame RL is represented as a one-latent-frame Wan run:
# num_latent_t: 1
# num_frames: 1
#
# The method trains the full transformer (no LoRA) and keeps an old-policy
# transformer plus a frozen reference transformer, matching the non-LoRA
# DiffusionNFT loss path.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
old:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
reference:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod
reward_fn:
pickscore: 1.0
clipscore: 1.0
sampling:
num_steps: 25
scheduler: flow_match_euler
trajectory: ode
flow_shift: inherit
validation:
every_steps: 10
num_steps: 40
num_prompts: 16
batch_size: 16
log_samples: true
seed: 42
# Null reuses training.data.data_path. Override this with a held-out
# preprocessed parquet path when one is available.
data_path:
# DiffusionNFT sd3_multi_reward on 4 GPUs resolves to per-GPU sample batch
# size 6, 48 sample batches per outer epoch, and grad accumulation 48.
sample_train_batch_size: 6
train_batch_size: 6
num_batches_per_epoch: 48
num_video_per_prompt: 24
num_inner_epochs: 1
timestep_fraction: 0.99
beta: 0.1
kl_beta: 0.0001
decay_type: 1
adv_mode: all
adv_clip_max: 5
max_grad_norm: 1.0
ema:
enabled: true
decay: 0.9
update_after_step: 0
validation: true
terminal_progress: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/pickscore_text_only_preprocessed
preprocessed_data_type: text_only
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 1
num_height: 448
num_width: 832
num_frames: 1
optimizer:
learning_rate: 3.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0001
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 48
checkpoint:
output_dir: outputs/wan2.1_diffusion_nft_pick_clip
training_state_checkpointing_steps: 30
checkpoints_total_limit: 3
tracker:
project_name: diffusion_nft_wan
run_name: wan2.1_diffusion_nft_pick_clip
model:
enable_gradient_checkpointing_type: full
pipeline:
flow_shift: 8
+9 -19
View File
@@ -78,26 +78,16 @@ print(f'{mj}.{mn}')"
}
if [ "${GPU_BACKEND}" = "CUDA" ]; then
# 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
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##*.}"
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
-5
View File
@@ -29,10 +29,6 @@ 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)
@@ -68,7 +64,6 @@ 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,11 +150,14 @@ 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(). 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.
# 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).
tile_buf: torch.Tensor | None = None
cache_tile_buf: bool = True
class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
@@ -172,7 +175,6 @@ 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
@@ -199,8 +201,7 @@ 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,
cache_tile_buf=cache_tile_buf)
untile_combined_index=untile_combined_index)
class VideoSparseAttentionImpl(AttentionImpl):
@@ -236,11 +237,6 @@ 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 -2
View File
@@ -1,6 +1,5 @@
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
@@ -15,5 +14,5 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
"MagiHumanVideoConfig", "StableAudioConfig"
]
-5
View File
@@ -14,11 +14,6 @@ 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,
-77
View File
@@ -1,77 +0,0 @@
# 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,8 +7,6 @@ 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
@@ -17,5 +15,5 @@ __all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
"StableAudioConditionerConfig", "T5GemmaEncoderConfig"
]
@@ -36,11 +36,6 @@ 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 = {
@@ -1,38 +0,0 @@
# 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
@@ -1,82 +0,0 @@
# 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,7 +6,6 @@ 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__ = [
@@ -20,5 +19,4 @@ __all__ = [
"LTX2VAEConfig",
"OobleckVAEArchConfig",
"OobleckVAEConfig",
"Flux2VAEConfig",
]
-58
View File
@@ -1,58 +0,0 @@
# 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
+4 -4
View File
@@ -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 (LucyEditDevConfig, SelfForcingWanT2V480PConfig, WanI2V480PConfig,
WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig)
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"SelfForcingWanT2V480PConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+1 -7
View File
@@ -35,11 +35,6 @@ 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
@@ -69,9 +64,8 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
# Wan2.2 task modifiers
# Wan2.2 TI2V parameters
ti2v_task: bool = False
lucy_edit_task: bool = False
boundary_ratio: float | None = None
# Compilation
-86
View File
@@ -1,86 +0,0 @@
# 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, ))
-138
View File
@@ -6,11 +6,9 @@ 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
@@ -122,142 +120,6 @@ 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
-7
View File
@@ -517,7 +517,6 @@ 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
@@ -568,8 +567,6 @@ 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,
@@ -672,7 +669,6 @@ 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:
@@ -719,9 +715,6 @@ 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 -38
View File
@@ -2,9 +2,8 @@
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, an
absolute VLM scorer (`videoscore2`), and a pairwise VLM judge
(`judge.third_person_separation`) — all behind a single registry-driven API.
comparisons, the full VBench suite, Physics-IQ, audio metrics, and a
VLM scorer behind a single registry-driven API.
## Install
@@ -160,7 +159,6 @@ 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
@@ -315,40 +313,6 @@ 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.
@@ -1,306 +0,0 @@
"""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))
-2
View File
@@ -85,8 +85,6 @@ 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"
-46
View File
@@ -219,15 +219,6 @@ 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,
@@ -294,8 +285,6 @@ 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
@@ -305,41 +294,6 @@ 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.
+2 -12
View File
@@ -5,7 +5,6 @@ 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):
@@ -22,27 +21,18 @@ 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,
quant_config=quant_config,
prefix=f"{prefix}.fc_in",
)
params_dtype=dtype)
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,
quant_config=quant_config,
prefix=f"{prefix}.fc_out")
self.fc_out = ReplicatedLinear(mlp_hidden_dim, output_dim, bias=bias, params_dtype=dtype)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x, _ = self.fc_in(x)
+4 -11
View File
@@ -52,7 +52,6 @@ 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
@@ -61,23 +60,17 @@ 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] if sequence_dim=2 else [B, S, H, D].
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
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
+25 -41
View File
@@ -26,7 +26,6 @@ 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
@@ -107,9 +106,7 @@ class WanSelfAttention(nn.Module):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
parallel_attention=False,
quant_config: QuantizationConfig | None = None,
prefix: str = "") -> None:
parallel_attention=False) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
@@ -121,10 +118,10 @@ class WanSelfAttention(nn.Module):
self.parallel_attention = parallel_attention
# layers
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.to_q = ReplicatedLinear(dim, dim)
self.to_k = ReplicatedLinear(dim, dim)
self.to_v = ReplicatedLinear(dim, dim)
self.to_out = ReplicatedLinear(dim, dim)
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()
@@ -197,15 +194,13 @@ class WanI2VCrossAttention(WanSelfAttention):
qk_norm=True,
eps=1e-6,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
| None = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends, quant_config=quant_config, prefix=prefix)
supported_attention_backends)
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.add_k_proj = ReplicatedLinear(dim, dim)
self.add_v_proj = ReplicatedLinear(dim, dim)
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()
@@ -251,17 +246,16 @@ 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, 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_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_out")
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
@@ -296,17 +290,13 @@ class WanTransformerBlock(nn.Module):
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.attn2")
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.attn2")
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
@@ -316,7 +306,7 @@ class WanTransformerBlock(nn.Module):
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
@@ -416,17 +406,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, 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.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.attn1 = DistributedAttention_VSA(
num_heads=num_heads,
head_size=dim // num_heads,
@@ -461,17 +451,13 @@ class WanTransformerBlock_VSA(nn.Module):
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.attn2")
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.attn2")
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
@@ -481,7 +467,7 @@ class WanTransformerBlock_VSA(nn.Module):
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
@@ -570,7 +556,6 @@ 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
@@ -609,7 +594,6 @@ 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)
])
-48
View File
@@ -1,48 +0,0 @@
# 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
-461
View File
@@ -1,461 +0,0 @@
# 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
+8 -66
View File
@@ -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, safe_open
from safetensors.torch import load_file as safetensors_load_file
from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
@@ -573,53 +573,13 @@ class TokenizerLoader(ComponentLoader):
# If parsing fails, fall through to AutoTokenizer below.
pass
# 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
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),
)
padding_side = None
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
try:
@@ -904,18 +864,6 @@ 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."""
@@ -951,12 +899,6 @@ 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
+2 -20
View File
@@ -332,7 +332,6 @@ 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
@@ -364,25 +363,8 @@ 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)
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
# 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(
-5
View File
@@ -42,7 +42,6 @@ _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 = {
@@ -70,9 +69,6 @@ _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] = {
@@ -94,7 +90,6 @@ _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"),
-532
View File
@@ -1,532 +0,0 @@
# 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)
-533
View File
@@ -1,533 +0,0 @@
# 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
@@ -1,7 +0,0 @@
# 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"]
@@ -1,17 +0,0 @@
# 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
@@ -1,138 +0,0 @@
# 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
@@ -1,95 +0,0 @@
# 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
@@ -1,161 +0,0 @@
# 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
@@ -1,111 +0,0 @@
# 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
@@ -1,75 +0,0 @@
# 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)
@@ -1,80 +0,0 @@
# 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
-19
View File
@@ -268,24 +268,6 @@ 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
# -------------------------------------------------------------------
@@ -359,7 +341,6 @@ 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,8 +380,6 @@ 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"
+2 -70
View File
@@ -47,23 +47,6 @@ 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.
@@ -94,29 +77,6 @@ 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:
"""
@@ -140,9 +100,7 @@ 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
# 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)
latents = self._denormalize_latents(latents)
# Decode latents
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
@@ -152,27 +110,7 @@ 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)
@@ -263,13 +201,7 @@ class DecodingStage(PipelineStage):
pipeline.add_module("vae", self.vae)
fastvideo_args.model_loaded["vae"] = True
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)
frames = batch.latents if fastvideo_args.output_type == "latent" else self.decode(batch.latents, fastvideo_args)
# decode trajectory latents if needed
if batch.return_trajectory_decoded:
+12 -110
View File
@@ -4,7 +4,6 @@ Denoising stage for diffusion pipelines.
"""
import inspect
import os
import weakref
from collections.abc import Iterable
from typing import Any
@@ -104,37 +103,7 @@ class DenoisingStage(PipelineStage):
# TODO(will): make the precision configurable for inference
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
# 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()
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
@@ -190,40 +159,13 @@ 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
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
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
@@ -270,34 +212,23 @@ 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) * (1.0 if _is_flux else 1000.0))
).to(target_dtype) * 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
@@ -378,25 +309,14 @@ class DenoisingStage(PipelineStage):
# Expand latents for V2V/I2V
latent_model_input = latents.to(target_dtype)
if batch.video_latent is not None:
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)
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 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:
if 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])
@@ -404,19 +324,7 @@ class DenoisingStage(PipelineStage):
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
# 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())
t_expand = t_expand.to(get_local_torch_device())
if use_meanflow:
if i == len(timesteps) - 1:
@@ -495,7 +403,6 @@ class DenoisingStage(PipelineStage):
**action_kwargs,
**camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
if batch.do_classifier_free_guidance:
@@ -537,7 +444,6 @@ class DenoisingStage(PipelineStage):
**action_kwargs,
**camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
_cfg_gate_fresh_uncond += 1
@@ -561,16 +467,12 @@ class DenoisingStage(PipelineStage):
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
if scheduler_fp32:
# Diffusers-style: fp32 Euler update outside autocast avoids BF16 drift.
# Compute the previous noisy sample
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[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)
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:
+3 -4
View File
@@ -629,10 +629,9 @@ class VideoVAEEncodingStage(ImageVAEEncodingStage):
encoder_output = self.vae.encode(video_condition)
generator = batch.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 generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
+2 -27
View File
@@ -57,10 +57,6 @@ 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
@@ -222,8 +218,7 @@ 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 isinstance(processed_text, str) and not processed_text.strip() and getattr(
encoder_config, "treat_empty_as_dot", False):
if not processed_text.strip() and getattr(encoder_config, "treat_empty_as_dot", False):
processed_text = "."
processed_texts.append(processed_text)
else:
@@ -240,27 +235,7 @@ class TextEncodingStage(PipelineStage):
tok = getattr(tokenizer, "tokenizer", tokenizer)
if encoder_config.is_chat_model:
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)
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(target_device)
else:
text_inputs = tok(processed_texts, **tok_kwargs).to(target_device)
+10 -70
View File
@@ -30,10 +30,6 @@ 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 (
@@ -44,7 +40,6 @@ 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,
@@ -329,45 +324,6 @@ 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,
@@ -782,18 +738,6 @@ 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,
@@ -896,19 +840,15 @@ def get_model_info(
if workload_type is None:
workload_type = WorkloadType.T2V
config_info = _get_config_info(model_path, raise_on_missing=True)
assert config_info is not None, "config_info must be resolved"
if override_pipeline_cls_name:
pipeline_name = override_pipeline_cls_name
logger.info("Using override pipeline class name %s", pipeline_name)
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
else:
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
else:
config = maybe_download_model_index(model_path)
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name")
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
if pipeline_name is None:
raise ValueError("Model config does not contain a _class_name attribute. "
@@ -917,6 +857,9 @@ 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(
@@ -977,12 +920,9 @@ 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,
-51
View File
@@ -1,51 +0,0 @@
# 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
@@ -1,38 +0,0 @@
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
@@ -1,49 +0,0 @@
# 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)
-531
View File
@@ -1,531 +0,0 @@
"""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())\"")
-29
View File
@@ -286,35 +286,6 @@ def run_train_framework_tests():
)
@app.function(gpu="L40S:1",
image=image,
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
],
volumes={"/root/data": model_vol})
def seed_grad_norm_references():
"""Record the per-method grad-norm reference for the **CI GPU (L40S only)**.
Phase 2 / 5a-ii one-off seeding entrypoint. Pinned to ``gpu="L40S:1"`` (the
Modal CI runner), so this function only seeds the ``L40S`` key in
``fastvideo/tests/train/methods/grad_norm_refs.json``.
``FASTVIDEO_GRADNORM_UPDATE=1`` makes ``check_grad_norm_regression`` record
the measured norm instead of asserting; ``-rs`` surfaces the recorded value
in the log so it can be copied into the JSON.
To seed any other device (e.g. our local Blackwell dev box → ``GB200``
key), run the same env-var + pytest invocation directly on that
workstation — see the module docstring of ``grad_norm_regression.py`` for
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)
@app.function(gpu="L40S:1",
image=image,
timeout=3600,
@@ -1,114 +0,0 @@
# 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"
),
},
)
@@ -67,8 +67,6 @@ class TestConstructor:
assert cb.num_frames is None
assert cb.sampling_timesteps is None
assert cb.output_dir is None
assert cb.offload_training_state is False
assert cb.unload_pipeline_after_validation is False
# Lazy fields not yet populated.
assert cb._pipeline is None
assert cb._sampling_param is None
@@ -85,16 +83,12 @@ class TestConstructor:
guidance_scale="4.5", # type: ignore[arg-type]
num_frames="77", # type: ignore[arg-type]
sampling_timesteps=["1000", "500"],
offload_training_state="1", # type: ignore[arg-type]
unload_pipeline_after_validation="false", # type: ignore[arg-type]
)
assert cb.every_steps == 50
assert cb.sampling_steps == [20, 40]
assert cb.guidance_scale == 4.5
assert cb.num_frames == 77
assert cb.sampling_timesteps == [1000, 500]
assert cb.offload_training_state is True
assert cb.unload_pipeline_after_validation is False
def test_pipeline_kwargs_collected(self) -> None:
cb = ValidationCallback(
@@ -1,10 +0,0 @@
{
"test_wan_causal_dfsft": {
"GB200": 2.9781,
"L40S": 3.2562
},
"test_wan_finetune": {
"GB200": 1.6486,
"L40S": 1.6467
}
}
@@ -1,158 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Layer-0 grad-norm regression for the per-method training smoke tests.
Phase 2 / 5a-ii: layers a device-keyed grad-norm check on top of the
finite/non-zero grad assertions established in 5a-i. After one
``single_train_step`` + ``backward``, the L2 norm of transformer block 0's
trainable gradients is compared against a reference value pinned per GPU in
``grad_norm_refs.json`` (next to this module).
Determinism: the harness seeds both the global RNG and the method's
``cuda_generator`` via ``method.on_train_start()`` (``training.data.seed`` in the
fixture), and the synthetic ``raw_batch`` is built *after* that call, so the
forward/backward is reproducible within bf16 reduction noise on a given GPU.
Why device-keyed: grad norms differ across GPU architectures (kernels,
accumulation order), so a single golden value can't cover every runner. The
JSON currently carries refs for the two GPUs we actually run on — ``L40S`` (CI)
and ``GB200`` (our Blackwell dev box; ``B200`` maps to the same key).
Seeding a reference for the current device:
- **CI / L40S** — invoke ``modal run`` against ``seed_grad_norm_references`` in
``fastvideo/tests/modal/pr_test.py`` (pinned to ``gpu="L40S:1"``), then copy
the recorded value from the log into ``grad_norm_refs.json``.
- **Local / non-L40S GPUs** — on that workstation::
FASTVIDEO_GRADNORM_UPDATE=1 \\
pytest fastvideo/tests/train/methods -vs -rs
The harness writes the measured norm into ``grad_norm_refs.json`` under the
device's key and skips the assertion for that run. Append a new substring
entry to ``_DEVICE_MAPPINGS`` first for any device not already listed.
"""
from __future__ import annotations
import json
import os
from pathlib import Path
import pytest
import torch
_REFS_PATH = Path(__file__).resolve().parent / "grad_norm_refs.json"
_UPDATE_ENV = "FASTVIDEO_GRADNORM_UPDATE"
# bf16 single-step smoke: catch gross breakage (wrong wiring, dead grads,
# scale regressions), not micro-drift from reduction nondeterminism.
_DEFAULT_RTOL = 0.10
# GPU-name substring -> reference key. First match wins. Only devices with
# seeded references in ``grad_norm_refs.json`` are listed here — to add a new
# GPU, append an entry, then seed the reference (see module docstring).
_DEVICE_MAPPINGS: tuple[tuple[str, str], ...] = (
("L40S", "L40S"),
("GB200", "GB200"),
("B200", "GB200"), # same Blackwell arch as GB200
)
def _device_name() -> str:
if not torch.cuda.is_available():
return "CPU"
return torch.cuda.get_device_name(0)
def resolve_device_key(device_name: str | None = None) -> str | None:
"""Map a CUDA device name to its reference key, or None if unsupported.
The substring match is case-insensitive so it survives driver/environment
differences in how ``torch.cuda.get_device_name`` capitalizes the model.
"""
name = device_name if device_name is not None else _device_name()
name_lower = name.lower()
for pattern, key in _DEVICE_MAPPINGS:
if pattern.lower() in name_lower:
return key
return None
def layer0_grad_norm(transformer) -> float:
"""Global L2 norm of transformer block 0's trainable gradients.
Block 0 is the reference surface 5a-i already isolates: its grad is the
*last* one produced during backprop, so a healthy value implies the whole
forward + chain-rule path is intact.
Accumulates the squared sums on the GPU and does a single CPU-GPU sync
(``.item()``) at the end, rather than one per parameter.
"""
blocks = getattr(transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0, (
"transformer is expected to expose a non-empty ``.blocks``")
grads = [
p.grad for p in blocks[0].parameters()
if p.requires_grad and p.grad is not None
]
if not grads:
return 0.0
sq_sum = torch.zeros((), device=grads[0].device, dtype=torch.float32)
for g in grads:
sq_sum += g.detach().float().pow(2).sum()
return sq_sum.sqrt().item()
def _load_refs() -> dict[str, dict[str, float]]:
if _REFS_PATH.exists():
return json.loads(_REFS_PATH.read_text(encoding="utf-8"))
return {}
def _save_refs(refs: dict[str, dict[str, float]]) -> None:
_REFS_PATH.write_text(
json.dumps(refs, indent=2, sort_keys=True) + "\n",
encoding="utf-8")
def check_grad_norm_regression(
test_name: str,
transformer,
*,
rtol: float = _DEFAULT_RTOL,
) -> None:
"""Assert block-0 grad norm matches the device-keyed reference within rtol.
- Skips when the current GPU has no reference (unsupported device, or not
yet seeded) so a new runner never hard-fails before its golden exists.
- With ``FASTVIDEO_GRADNORM_UPDATE=1`` records/updates the reference for the
current device instead of asserting.
"""
norm = layer0_grad_norm(transformer)
device_key = resolve_device_key()
if os.environ.get(_UPDATE_ENV) == "1":
if device_key is None:
pytest.skip(
f"{_UPDATE_ENV}=1 but GPU '{_device_name()}' has no reference "
"key; add it to _DEVICE_MAPPINGS first")
refs = _load_refs()
refs.setdefault(test_name, {})[device_key] = round(norm, 4)
_save_refs(refs)
pytest.skip(
f"recorded grad-norm reference {test_name}[{device_key}] = "
f"{norm:.4f} (assertion skipped under {_UPDATE_ENV}=1)")
ref = _load_refs().get(test_name, {}).get(device_key) \
if device_key is not None else None
if ref is None:
pytest.skip(
f"no grad-norm reference for {test_name} on '{_device_name()}' "
f"(device_key={device_key}); run with {_UPDATE_ENV}=1 to seed it")
rel = abs(norm - ref) / (abs(ref) + 1e-12)
assert rel <= rtol, (
f"{test_name}[{device_key}] grad-norm regression: got {norm:.4f}, "
f"reference {ref:.4f}, relative error {rel:.3%} exceeds rtol "
f"{rtol:.0%}. If this is an intentional change, refresh the reference "
f"with {_UPDATE_ENV}=1 and explain why in the PR.")
@@ -28,8 +28,6 @@ from fastvideo.train.methods.fine_tuning.dfsft import (
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -124,7 +122,3 @@ def test_wan_causal_dfsft_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
@@ -36,8 +36,6 @@ from fastvideo.train.methods.fine_tuning.finetune import (
from fastvideo.train.models.wan import WanModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -141,7 +139,3 @@ def test_wan_finetune_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_finetune", model.transformer)
+9 -186
View File
@@ -68,8 +68,6 @@ class ValidationCallback(Callback):
num_frames: int | None = None,
output_dir: str | None = None,
sampling_timesteps: list[int] | None = None,
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
**pipeline_kwargs: Any,
) -> None:
self.pipeline_target = str(pipeline_target)
@@ -80,8 +78,6 @@ class ValidationCallback(Callback):
self.num_frames = (int(num_frames) if num_frames is not None else None)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None)
self.offload_training_state = self._coerce_bool(offload_training_state)
self.unload_pipeline_after_validation = self._coerce_bool(unload_pipeline_after_validation)
self.pipeline_kwargs = dict(pipeline_kwargs)
# Set after on_train_start.
@@ -92,12 +88,6 @@ class ValidationCallback(Callback):
self.validation_random_generator: (torch.Generator | None) = None
self.seed: int = 0
@staticmethod
def _coerce_bool(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return bool(value)
# ----------------------------------------------------------
# Callback hooks
# ----------------------------------------------------------
@@ -150,183 +140,16 @@ class ValidationCallback(Callback):
) -> None:
transformer = method.student.transformer
try:
with self._validation_memory_context(
method,
validation_transformer=transformer,
):
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
method,
step,
t,
)
finally:
if self.unload_pipeline_after_validation:
self._clear_pipeline_cache()
@contextlib.contextmanager
def _validation_memory_context(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
):
if not self.offload_training_state:
yield
return
optimizer_tensor_records: list[tuple[Any, Any, torch.device]] = []
module_records: list[tuple[str, torch.nn.Module, torch.device]] = []
try:
self._offload_optimizer_states_to_cpu(
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
method,
optimizer_tensor_records,
step,
t,
)
self._offload_inactive_role_modules_to_cpu(
method,
validation_transformer=validation_transformer,
module_records=module_records,
)
self._empty_cuda_cache()
yield
finally:
self._restore_inactive_role_modules(module_records)
self._restore_optimizer_states(optimizer_tensor_records)
self._empty_cuda_cache()
def _offload_optimizer_states_to_cpu(
self,
method: TrainingMethod,
records: list[tuple[Any, Any, torch.device]],
) -> None:
optimizers = getattr(method, "_optimizer_dict", {})
if not optimizers:
return
moved = 0
for optimizer in optimizers.values():
state = getattr(optimizer, "state", None)
if not isinstance(state, dict):
continue
for param_state in state.values():
moved += self._offload_tensor_container_to_cpu(
param_state,
records,
)
if moved:
logger.info(
"Offloaded %d optimizer state tensors to CPU for validation.",
moved,
)
def _offload_tensor_container_to_cpu(
self,
obj: Any,
records: list[tuple[Any, Any, torch.device]],
) -> int:
moved = 0
if isinstance(obj, dict):
for key, value in list(obj.items()):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, key, value.device))
obj[key] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
if isinstance(obj, list):
for idx, value in enumerate(list(obj)):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, idx, value.device))
obj[idx] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
def _restore_optimizer_states(
self,
records: list[tuple[Any, Any, torch.device]],
) -> None:
for container, key, device in reversed(records):
value = container[key]
if torch.is_tensor(value):
container[key] = value.to(device=device)
if records:
logger.info(
"Restored %d optimizer state tensors after validation.",
len(records),
)
def _offload_inactive_role_modules_to_cpu(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
role_models = getattr(method, "_role_models", {})
if not isinstance(role_models, dict):
return
for role, model in role_models.items():
module = getattr(model, "transformer", None)
if not isinstance(module, torch.nn.Module):
continue
if module is validation_transformer:
continue
device = self._first_cuda_tensor_device(module)
if device is None:
continue
try:
module.to("cpu")
except Exception as exc:
logger.warning(
"Could not offload role %r transformer to CPU before validation: %s",
role,
exc,
)
continue
module_records.append((str(role), module, device))
logger.info(
"Offloaded role %r transformer from %s to CPU for validation.",
role,
device,
)
def _restore_inactive_role_modules(
self,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
for role, module, device in reversed(module_records):
module.to(device)
logger.info(
"Restored role %r transformer to %s after validation.",
role,
device,
)
@staticmethod
def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None:
for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)):
device = getattr(tensor, "device", None)
if isinstance(device, torch.device) and device.type == "cuda":
return device
return None
def _clear_pipeline_cache(self) -> None:
self._pipeline = None
self._pipeline_key = None
self._empty_cuda_cache()
@staticmethod
def _empty_cuda_cache() -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _find_ema_callback(self) -> Any | None:
"""Find the EMA callback in the callback dict."""
@@ -470,7 +293,6 @@ class ValidationCallback(Callback):
}
if flow_shift is not None:
kwargs["flow_shift"] = float(flow_shift)
kwargs.update(self.pipeline_kwargs)
self._pipeline = PipelineCls.from_pretrained(
tc.model_path,
@@ -540,6 +362,7 @@ class ValidationCallback(Callback):
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=self.validation_random_generator,
n_tokens=n_tokens,
eta=0.0,
-40
View File
@@ -197,46 +197,6 @@ class TrainingMethod(torch.nn.Module, ABC):
# -- Shared hooks (override in subclasses as needed) --
def manages_optimization(self) -> bool:
"""Whether the method owns backward/optimizer stepping internally.
Most methods return loss tensors and let :class:`Trainer` handle
gradient accumulation, callbacks, optimizer stepping, and scheduler
stepping. RL-style methods such as DiffusionNFT need to preserve their
own sample-then-inner-train loop, so they can provide a specific
``managed_train_step``.
"""
return False
def managed_train_step(
self,
data_stream: Any,
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
"""Run one method-managed step.
Subclasses that return ``True`` from :meth:`manages_optimization`
should override this. The fallback consumes one dataloader batch and
delegates to ``single_train_step`` so tests can exercise the hook with
tiny fake methods.
"""
return self.single_train_step(next(data_stream), iteration)
def on_validation_begin(self, iteration: int = 0) -> dict[str, LogScalar]:
"""Run method-owned validation, if any.
Pipeline-style validation should remain in callbacks. Methods that
intentionally avoid inference pipelines, such as RL methods with their
own sampler/reward loop, can override this hook and return metrics for
the trainer to log at ``iteration``.
"""
del iteration
return {}
def get_grad_clip_targets(
self,
iteration: int,
@@ -55,8 +55,6 @@ class DMD2Method(TrainingMethod):
raise ValueError("DMD2Method requires critic to be trainable")
self._cfg_uncond = self._parse_cfg_uncond()
self._rollout_mode = self._parse_rollout_mode()
self._validate_preprocessed_data_type()
self._configure_student_negative_conditioning()
self._denoising_step_list: torch.Tensor | None = (None)
# Initialize preprocessors on student.
@@ -208,13 +206,6 @@ class DMD2Method(TrainingMethod):
return targets
def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]:
"""Parse how DMD2 obtains the latent point used for rollout.
``simulate`` starts from fresh noise and lets the student create an
artificial latent trajectory, so it can run with text-only data.
``data_latent`` starts from preprocessed VAE latents and perturbs them
at a sampled denoising timestep.
"""
raw = self.method_config.get("rollout_mode", None)
if raw is None:
raise ValueError("method_config.rollout_mode must be set "
@@ -232,34 +223,6 @@ class DMD2Method(TrainingMethod):
"{simulate, data_latent}, got "
f"{raw!r}")
def _validate_preprocessed_data_type(self) -> None:
data_type = str(getattr(
self.training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
if data_type == "text_only" and self._rollout_mode != "simulate":
raise ValueError("training.data.preprocessed_data_type='text_only' "
"requires method.rollout_mode='simulate'; "
"data_latent rollout requires vae_latent data.")
def _uses_negative_prompt_conditioning(self) -> bool:
if self._cfg_uncond is None:
return True
text_policy = self._cfg_uncond.get("text", None)
if text_policy is None:
return True
return str(text_policy).strip().lower() == "negative_prompt"
def _configure_student_negative_conditioning(self) -> None:
setter = getattr(
self.student,
"set_requires_negative_conditioning",
None,
)
if setter is not None:
setter(self._uses_negative_prompt_conditioning())
def _parse_cfg_uncond(self, ) -> dict[str, Any] | None:
raw = self.method_config.get("cfg_uncond", None)
if raw is None:
-6
View File
@@ -1,6 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""RL training methods."""
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
__all__ = ["DiffusionNFTMethod"]
@@ -1,30 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable RL training primitives."""
from fastvideo.train.methods.rl.common.sampling import (
DiffusionSampler,
SamplingConfig,
SamplingResult,
)
from fastvideo.train.methods.rl.common.prompt_sampling import (
KRepeatSample,
distributed_k_repeat_indices,
)
from fastvideo.train.methods.rl.common.validation import (
RLValidationConfig,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
__all__ = [
"DiffusionSampler",
"KRepeatSample",
"RLValidationConfig",
"SamplingConfig",
"SamplingResult",
"distributed_k_repeat_indices",
"media_to_video_array",
"validation_caption",
"validation_shard_indices",
]
@@ -1,74 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Prompt-row sampling helpers for online RL methods.
This module chooses and repeats dataset prompt rows across ranks for RL training
batches. Here, "sampling" means selection, not generator sampling.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True, slots=True)
class KRepeatSample:
"""Local prompt indices for one distributed K-repeat sampling batch."""
local_indices: list[int]
unique_prompt_count: int
def distributed_k_repeat_indices(
*,
dataset_length: int,
batch_size: int,
repeats_per_prompt: int,
world_size: int,
rank: int,
seed: int,
) -> KRepeatSample:
"""Mirror DiffusionNFT's distributed K-repeat prompt sampler.
Adapted from DiffusionNFT's
``scripts/train_nft_sd3.py::DistributedKRepeatSampler``.
"""
dataset_length = int(dataset_length)
batch_size = int(batch_size)
repeats_per_prompt = int(repeats_per_prompt)
world_size = int(world_size)
rank = int(rank)
if dataset_length <= 0:
raise ValueError("dataset_length must be positive")
if batch_size <= 0:
raise ValueError("batch_size must be positive")
if repeats_per_prompt <= 0:
raise ValueError("repeats_per_prompt must be positive")
if world_size <= 0:
raise ValueError("world_size must be positive")
if rank < 0 or rank >= world_size:
raise ValueError(f"rank must be in [0, {world_size}), got {rank}")
total_samples = world_size * batch_size
if total_samples % repeats_per_prompt != 0:
raise ValueError("world_size * batch_size must be divisible by repeats_per_prompt "
f"({world_size} * {batch_size} vs {repeats_per_prompt})")
unique_prompt_count = total_samples // repeats_per_prompt
if unique_prompt_count > dataset_length:
raise ValueError("K-repeat sampling needs at least as many rows as unique prompts "
f"per sampling batch ({dataset_length} < {unique_prompt_count})")
generator = torch.Generator()
generator.manual_seed(int(seed))
indices = torch.randperm(dataset_length, generator=generator)[:unique_prompt_count].tolist()
repeated_indices = [idx for idx in indices for _ in range(repeats_per_prompt)]
shuffled_order = torch.randperm(len(repeated_indices), generator=generator).tolist()
shuffled_samples = [int(repeated_indices[idx]) for idx in shuffled_order]
start = rank * batch_size
end = start + batch_size
return KRepeatSample(
local_indices=shuffled_samples[start:end],
unique_prompt_count=unique_prompt_count,
)
@@ -1,223 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Configurable diffusion samplers for RL training methods."""
from __future__ import annotations
import copy
from dataclasses import dataclass
from typing import Any, Literal
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import ModelBase
SchedulerName = Literal["flow_match_euler", "model_default"]
TrajectoryName = Literal["ode", "sde_reflow"]
@dataclass(slots=True)
class SamplingConfig:
"""YAML-backed sampling knobs shared by RL methods."""
num_steps: int = 25
scheduler: SchedulerName = "model_default"
trajectory: TrajectoryName = "ode"
flow_shift: float | None = None
timesteps: list[float] | None = None
sigmas: list[float] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> SamplingConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.sampling must be a mapping, got {type(raw).__name__}")
supported_keys = {
"flow_shift",
"num_steps",
"scheduler",
"sigmas",
"timesteps",
"trajectory",
}
unsupported_keys = sorted(set(raw) - supported_keys)
if unsupported_keys:
raise ValueError(f"Unsupported method.sampling key(s): {unsupported_keys}. "
f"Supported keys: {sorted(supported_keys)}")
scheduler = str(raw.get("scheduler", "model_default") or "model_default").strip().lower()
if scheduler not in {"flow_match_euler", "model_default"}:
raise ValueError("method.sampling.scheduler must be one of "
"{flow_match_euler, model_default}, got "
f"{raw.get('scheduler')!r}")
trajectory = str(raw.get("trajectory", "ode") or "ode").strip().lower()
if trajectory not in {"ode", "sde_reflow"}:
raise ValueError("method.sampling.trajectory must be one of "
"{ode, sde_reflow}, got "
f"{raw.get('trajectory')!r}")
timesteps = raw.get("timesteps")
sigmas = raw.get("sigmas")
if timesteps is not None:
if not isinstance(timesteps, list) or not timesteps:
raise ValueError("method.sampling.timesteps must be a non-empty list when set")
timesteps = [float(t) for t in timesteps]
if sigmas is not None:
if not isinstance(sigmas, list) or not sigmas:
raise ValueError("method.sampling.sigmas must be a non-empty list when set")
sigmas = [float(s) for s in sigmas]
if timesteps is not None and sigmas is not None and len(timesteps) != len(sigmas):
raise ValueError("method.sampling.timesteps and method.sampling.sigmas must have the same length")
num_steps = int(raw.get("num_steps", 25) or 25)
if num_steps <= 0:
raise ValueError("method.sampling.num_steps must be positive")
return cls(
num_steps=num_steps,
scheduler=scheduler, # type: ignore[arg-type]
trajectory=trajectory, # type: ignore[arg-type]
flow_shift=(None if raw.get("flow_shift", None) in (None, "inherit") else float(raw["flow_shift"])),
timesteps=timesteps,
sigmas=sigmas,
)
@dataclass(slots=True)
class SamplingResult:
latents: torch.Tensor
timesteps: torch.Tensor
sigmas: torch.Tensor
class DiffusionSampler:
"""Thin model/scheduler sampler used by RL methods.
This intentionally does not call FastVideo's full inference pipelines.
RL training needs a reusable sampling primitive that works with
``ModelBase`` wrappers and scheduler math without binding a method to
model-family pipeline classes such as ``WanDMDPipeline``.
"""
def __init__(self, config: SamplingConfig) -> None:
self.config = config
@torch.no_grad()
def sample(
self,
model: ModelBase,
batch: TrainingBatch,
*,
generator: torch.Generator | None,
) -> SamplingResult:
latents = batch.latents
if latents is None:
raise RuntimeError("TrainingBatch.latents is required for RL sampling")
current = torch.randn(
latents.shape,
device=latents.device,
dtype=latents.dtype,
generator=generator,
)
scheduler = self._prepare_scheduler(model, current.device)
timesteps = scheduler.timesteps.to(device=current.device)
sigmas = scheduler.sigmas.to(device=current.device)
original_timesteps = batch.timesteps
try:
if self.config.trajectory == "ode":
pred_clean = current
for timestep in timesteps:
model_timestep = self._model_timestep(timestep, current)
batch.timesteps = model_timestep
pred_noise = model.predict_noise(
current,
model_timestep,
batch,
conditional=True,
attn_kind="dense",
)
current = scheduler.step(
pred_noise.flatten(0, 1),
timestep,
current.flatten(0, 1),
return_dict=False,
)[0].unflatten(0, pred_noise.shape[:2])
pred_clean = current
return SamplingResult(latents=pred_clean, timesteps=timesteps, sigmas=sigmas)
return SamplingResult(
latents=self._sample_sde_reflow(
model,
batch,
current,
timesteps,
generator=generator,
),
timesteps=timesteps,
sigmas=sigmas,
)
finally:
batch.timesteps = original_timesteps
def _prepare_scheduler(
self,
model: ModelBase,
device: torch.device,
) -> Any:
if self.config.scheduler == "flow_match_euler":
shift = self.config.flow_shift
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(shift))
else:
scheduler = copy.deepcopy(model.noise_scheduler)
kwargs: dict[str, Any] = {"device": device}
if self.config.timesteps is not None:
kwargs["timesteps"] = self.config.timesteps
kwargs["num_inference_steps"] = len(self.config.timesteps)
if self.config.sigmas is not None:
kwargs["sigmas"] = self.config.sigmas
kwargs["num_inference_steps"] = len(self.config.sigmas)
if "num_inference_steps" not in kwargs:
kwargs["num_inference_steps"] = self.config.num_steps
scheduler.set_timesteps(**kwargs)
return scheduler
def _sample_sde_reflow(
self,
model: ModelBase,
batch: TrainingBatch,
current: torch.Tensor,
timesteps: torch.Tensor,
*,
generator: torch.Generator | None,
) -> torch.Tensor:
pred_clean = current
for step_idx, timestep in enumerate(timesteps):
timestep_tensor = self._model_timestep(timestep, current)
batch.timesteps = timestep_tensor
pred_clean = model.predict_x0(
current,
timestep_tensor,
batch,
conditional=True,
attn_kind="dense",
)
if step_idx < len(timesteps) - 1:
next_timestep = timesteps[step_idx + 1].reshape(1).to(device=current.device)
noise = torch.randn(
pred_clean.shape,
device=pred_clean.device,
dtype=pred_clean.dtype,
generator=generator,
)
current = model.add_noise(pred_clean, noise, next_timestep)
return pred_clean
@staticmethod
def _model_timestep(
timestep: torch.Tensor,
current: torch.Tensor,
) -> torch.Tensor:
return timestep.reshape(1).to(device=current.device).expand(current.shape[0]).contiguous()
@@ -1,81 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared validation helpers for RL training methods."""
from __future__ import annotations
from dataclasses import dataclass
import math
from typing import Any
import torch
@dataclass(slots=True)
class RLValidationConfig:
every_steps: int = 0
num_steps: int = 40 # Reference DiffusionNFT sampling num steps for best visual quality
num_prompts: int = 16
batch_size: int = 16
log_samples: bool = True
seed: int = 42
data_path: str | None = None
sampling: dict[str, Any] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> RLValidationConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.validation must be a mapping, got {type(raw).__name__}")
data_path = raw.get("data_path", None)
sampling = raw.get("sampling", None)
if sampling is not None and not isinstance(sampling, dict):
raise ValueError(f"method.validation.sampling must be a mapping, got {type(sampling).__name__}")
return cls(
every_steps=max(0, int(raw.get("every_steps", 0) or 0)),
num_steps=max(1, int(raw.get("num_steps", 40) or 40)),
num_prompts=max(1, int(raw.get("num_prompts", 16) or 16)),
batch_size=max(1, int(raw.get("batch_size", 16) or 16)),
log_samples=bool(raw.get("log_samples", True)),
seed=int(raw.get("seed", 42) or 42),
data_path=(None if data_path in (None, "") else str(data_path)),
sampling=(dict(sampling) if sampling is not None else None),
)
def validation_shard_indices(
num_prompts: int,
*,
rank: int,
world_size: int,
) -> list[tuple[int, bool]]:
"""Return fixed validation prompt indices for one distributed rank."""
num_prompts = max(1, int(num_prompts))
world_size = max(1, int(world_size))
per_rank = int(math.ceil(num_prompts / world_size))
padded_total = per_rank * world_size
return [((idx % num_prompts), idx < num_prompts) for idx in range(rank, padded_total, world_size)]
def validation_caption(
prompt: str,
rewards: dict[str, float],
) -> str:
reward_parts = [f"{key}: {float(rewards[key]):.4f}" for key in sorted(rewards)]
return f"{' | '.join(reward_parts)} | {prompt[:1000]}"
def media_to_video_array(media: torch.Tensor) -> Any:
"""Convert decoded media to a tracker video array.
Accepts ``[C, T, H, W]`` tensors. ``[C, H, W]`` tensors are treated as
``T=1`` media. Output follows the existing tracker convention used
elsewhere in FastVideo: ``[T, C, H, W]`` uint8.
"""
if media.ndim == 3:
media = media.unsqueeze(1)
if media.ndim != 4:
raise ValueError("media must have shape [C, T, H, W] or [C, H, W], "
f"got {tuple(media.shape)}")
video = (media.detach().float().clamp(0, 1) * 255).round().to(torch.uint8)
return video.permute(1, 0, 2, 3).contiguous().cpu().numpy()
File diff suppressed because it is too large Load Diff
@@ -1,37 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
select_first_frame,
)
def build_multi_reward_scorer(
reward_weights,
*,
device="cuda",
scorers: dict[str, RewardScorer] | None = None,
) -> MultiRewardScorer:
available: dict[str, RewardScorer] = dict(scorers or {})
if not available:
available = {
"pickscore": PickScoreScorer(device=device),
"clipscore": ClipScoreScorer(device=device),
}
return MultiRewardScorer(reward_weights, scorers=available)
__all__ = [
"ClipScoreScorer",
"MultiRewardScorer",
"PickScoreScorer",
"RewardScorer",
"build_multi_reward_scorer",
"select_first_frame",
]
@@ -1,130 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Frame-based reward scorers used by RL training methods."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from PIL import Image
import torch
from fastvideo.train.methods.rl.rewards.media import select_first_frame
class PickScoreScorer(torch.nn.Module):
"""PickScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/pickscore_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
from transformers import AutoModel, AutoProcessor
processor_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
model_path = "yuvalkirstain/PickScore_v1"
self.device = torch.device(device)
self.dtype = dtype
self.processor = AutoProcessor.from_pretrained(processor_path)
self.model = AutoModel.from_pretrained(model_path).eval().to(self.device)
self.model = self.model.to(dtype=dtype)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media)
frame_np = (frame_tensor.detach().float().clamp(0, 1) * 255).round()
frame_np = frame_np.to(torch.uint8).cpu().numpy().transpose(0, 2, 3, 1)
pil_frames = [Image.fromarray(frame) for frame in frame_np]
frame_inputs = self.processor(
images=pil_frames,
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
frame_inputs = {k: v.to(device=self.device) for k, v in frame_inputs.items()}
text_inputs = self.processor(
text=list(prompts),
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
text_inputs = {k: v.to(device=self.device) for k, v in text_inputs.items()}
text_embs = self.model.get_text_features(**text_inputs)
text_embs = text_embs / text_embs.norm(p=2, dim=-1, keepdim=True)
frame_embs = self.model.get_image_features(**frame_inputs)
frame_embs = frame_embs / frame_embs.norm(p=2, dim=-1, keepdim=True)
scores = self.model.logit_scale.exp() * (text_embs @ frame_embs.T)
return scores.diag().float() / 26.0
class ClipScoreScorer(torch.nn.Module):
"""CLIPScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/clip_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
) -> None:
super().__init__()
import torch.nn as nn
import torchvision.transforms as T
from transformers import CLIPModel, CLIPProcessor
def get_size(size: Any) -> Any:
if isinstance(size, int):
return (size, size)
if isinstance(size, Mapping) and "height" in size and "width" in size:
return (size["height"], size["width"])
if isinstance(size, Mapping) and "shortest_edge" in size:
return size["shortest_edge"]
raise ValueError(f"Invalid processor size: {size!r}")
def get_frame_transform(processor: Any) -> torch.nn.Module:
config = processor.to_dict()
resize = T.Resize(get_size(config.get("size"))) if config.get("do_resize") else nn.Identity()
crop = T.CenterCrop(get_size(config.get("crop_size"))) if config.get("do_center_crop") else nn.Identity()
normalize = (T.Normalize(mean=processor.image_mean, std=processor.image_std)
if config.get("do_normalize") else nn.Identity())
return T.Compose([resize, crop, normalize])
self.device = torch.device(device)
self.model = CLIPModel.from_pretrained("openai/clip-vit-large-patch14").to(self.device).eval()
self.processor = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
self.transform = get_frame_transform(self.processor.image_processor)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media).detach().float().clamp(0, 1)
texts = self.processor(
text=list(prompts),
padding="max_length",
truncation=True,
return_tensors="pt",
).to(self.device)
pixels = self.transform(frame_tensor).to(device=self.device, dtype=frame_tensor.dtype)
outputs = self.model(pixel_values=pixels, **texts)
return outputs.logits_per_image.diagonal().float() / 100.0
@@ -1,73 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Generic media reward composition utilities."""
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import torch
RewardScorer = Callable[[torch.Tensor, Sequence[str]], torch.Tensor]
def select_first_frame(media: torch.Tensor) -> torch.Tensor:
"""Return first-frame media as ``[B, C, H, W]``.
This is a helper for reward models that are intrinsically frame-based
(for example PickScore and CLIPScore). Video-aware rewards should inspect
the full ``[B, C, T, H, W]`` tensor themselves.
"""
if not torch.is_tensor(media):
raise TypeError(f"media must be a torch.Tensor, got {type(media).__name__}")
if media.ndim == 5:
return media[:, :, 0]
if media.ndim == 4:
return media
raise ValueError("media must have shape [B, C, H, W] or [B, C, T, H, W], "
f"got {tuple(media.shape)}")
class MultiRewardScorer:
"""Weighted sum of reusable media reward scorers.
Mirrors DiffusionNFT's ``flow_grpo/rewards.py::multi_score`` behavior,
while leaving frame selection to each concrete reward.
"""
def __init__(
self,
reward_weights: Mapping[str, float],
*,
scorers: Mapping[str, RewardScorer],
) -> None:
self.reward_weights = {str(k): float(v) for k, v in reward_weights.items()}
if not self.reward_weights:
raise ValueError("reward_weights must contain at least one reward")
self.scorers = dict(scorers)
unsupported = sorted(set(self.reward_weights) - set(self.scorers))
if unsupported:
raise ValueError(f"Unsupported reward(s): {unsupported}. "
f"Available rewards: {sorted(self.scorers)}")
@torch.no_grad()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> dict[str, torch.Tensor]:
prompt_count = len(prompts)
if media.shape[0] != prompt_count:
raise ValueError(f"media batch size ({media.shape[0]}) must match prompt count ({prompt_count})")
total: torch.Tensor | None = None
details: dict[str, torch.Tensor] = {}
for name, weight in self.reward_weights.items():
scores = self.scorers[name](media, prompts).detach().float()
if scores.ndim != 1 or int(scores.shape[0]) != prompt_count:
raise ValueError(f"Reward {name!r} must return shape [{prompt_count}], got {tuple(scores.shape)}")
details[name] = scores
weighted = scores * float(weight)
total = weighted if total is None else total.to(weighted.device) + weighted
assert total is not None
details["avg"] = total
return details
-11
View File
@@ -91,17 +91,6 @@ class ModelBase(ABC):
def on_train_start(self) -> None: # noqa: B027
"""Called once before the training loop begins."""
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
"""Decode ``[B, T, C, H, W]`` latents to ``[B, C, T, H, W]`` media.
RL reward methods call this hook instead of reaching into
model-specific VAE normalization details.
"""
raise NotImplementedError(f"{type(self).__name__} does not implement decode_latents()")
# ------------------------------------------------------------------
# Timestep helpers
# ------------------------------------------------------------------
+5 -48
View File
@@ -101,7 +101,6 @@ class WanModel(ModelBase):
self.negative_prompt_embeds: (torch.Tensor | None) = None
self.negative_prompt_attention_mask: (torch.Tensor | None) = None
self._requires_negative_conditioning = True
# Timestep mechanics.
self.timestep_shift: float = float(flow_shift)
@@ -161,31 +160,17 @@ class WanModel(ModelBase):
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_t2v,
pyarrow_schema_text_only,
)
pyarrow_schema_t2v, )
from fastvideo.train.utils.dataloader import (
build_parquet_t2v_train_dataloader, )
preprocessed_data_type = str(getattr(
training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
parquet_schema = pyarrow_schema_t2v
if preprocessed_data_type == "text_only":
parquet_schema = pyarrow_schema_text_only
elif preprocessed_data_type != "t2v":
raise ValueError("Unsupported Wan preprocessed_data_type: "
f"{preprocessed_data_type!r}")
text_len = (
training_config.pipeline_config.text_encoder_configs[ # type: ignore[union-attr]
0].arch_config.text_len)
self.dataloader = build_parquet_t2v_train_dataloader(
training_config.data,
text_len=int(text_len),
parquet_schema=parquet_schema,
parquet_schema=pyarrow_schema_t2v,
)
self.start_step = 0
@@ -193,9 +178,6 @@ class WanModel(ModelBase):
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def set_requires_negative_conditioning(self, requires: bool) -> None:
self._requires_negative_conditioning = bool(requires)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
@@ -205,25 +187,7 @@ class WanModel(ModelBase):
return timestep.clamp(self.min_timestep, self.max_timestep)
def on_train_start(self) -> None:
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@torch.no_grad()
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
if self.vae is None:
raise RuntimeError("Wan VAE is not initialized")
latents = latents_b_t_c_h_w.permute(0, 2, 1, 3, 4).float()
if bool(getattr(self.vae, "handles_latent_denorm", False)):
denorm = latents
else:
mean = torch.tensor(self.vae.latents_mean, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(self.vae.latents_std, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
denorm = latents * std + mean
media = self.vae.to(latents.device).decode(denorm)
return (media / 2 + 0.5).clamp(0, 1)
self.ensure_negative_conditioning()
# ------------------------------------------------------------------
# Runtime primitives
@@ -236,8 +200,7 @@ class WanModel(ModelBase):
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
self.ensure_negative_conditioning()
assert self.training_config is not None
tc = self.training_config
@@ -322,7 +285,7 @@ class WanModel(ModelBase):
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = self._get_training_dtype()
dtype = noisy_latents.dtype
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
@@ -338,11 +301,6 @@ class WanModel(ModelBase):
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
# Keep Wan training autocast tied to the model's training dtype, not
# to caller-created intermediates that may accidentally be fp32.
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
@@ -443,7 +401,6 @@ 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):
+1 -3
View File
@@ -128,9 +128,7 @@ class WanCausalModel(WanModel, CausalModelBase):
}
device_type = self.device.type
dtype = self._get_training_dtype()
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
dtype = noisy_latents.dtype
if conditional:
text_dict = batch.conditional_dict

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