Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4bcaf9bef8 | ||
|
|
88a5a933ab | ||
|
|
c591d6d2a6 | ||
|
|
65dff806a8 | ||
|
|
b85f0f4c2a | ||
|
|
76c62d7a00 | ||
|
|
f6e65ff668 | ||
|
|
c220aa8000 | ||
|
|
4713fc17ed | ||
|
|
5789955bbe | ||
|
|
2ad84a3b78 | ||
|
|
12d699cd78 | ||
|
|
34f14ded21 | ||
|
|
71d1ab411f | ||
|
|
805e487773 | ||
|
|
8803b4547e | ||
|
|
3b3806b3f6 | ||
|
|
38d962e89d | ||
|
|
3966a365d0 | ||
|
|
d73fd14af0 |
+186
-2
@@ -9,11 +9,183 @@ notify:
|
||||
- github_commit_status:
|
||||
context: "full-suite-passed"
|
||||
if: build.env("TEST_SCOPE") == "full"
|
||||
- github_commit_status:
|
||||
context: "direct-test-completed"
|
||||
if: build.env("TEST_SCOPE") == "direct"
|
||||
|
||||
steps:
|
||||
# ============================================================
|
||||
- label: ":dart: Direct Test (${TEST_TYPE})"
|
||||
if: build.env("TEST_SCOPE") == "direct"
|
||||
# Direct test: triggered by /test <name> slash command.
|
||||
# Labels match fastcheck/full-suite counterparts so the GitHub
|
||||
# check status overwrites the original failed check.
|
||||
# Only ONE step executes per build (gated by TEST_TYPE).
|
||||
# ============================================================
|
||||
|
||||
# --- Fastcheck-scope direct tests ---
|
||||
- label: ":microscope: Encoder Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "encoder"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: VAE Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "vae"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Transformer Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "transformer"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Kernel Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "kernel_tests"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Unit Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "unit_test"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
# --- Full-suite-scope direct tests ---
|
||||
- label: ":bar_chart: SSIM Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "ssim"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Inference Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_lora"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Distillation DMD Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "distillation_dmd"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Self-Forcing Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "self_forcing"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_lora"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests VSA"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_vsa"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Inference Tests VMoBA"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_vmoba"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Performance Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "performance"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: API Server Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "api_server"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
@@ -135,6 +307,10 @@ steps:
|
||||
label: ":bar_chart: SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
@@ -195,6 +371,10 @@ steps:
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
@@ -207,6 +387,10 @@ steps:
|
||||
label: ":test_tube: Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
|
||||
+6
-11
@@ -4,8 +4,10 @@ merge_protections:
|
||||
- base = main
|
||||
success_conditions:
|
||||
- "title~=(?i)^\\[(feat|feature|bugfix|fix|refactor|perf|ci|doc|docs|misc|chore|kernel|new.?model)\\]"
|
||||
- "#approved-reviews-by>=1"
|
||||
- check-success~=pre-commit
|
||||
- check-success=fastcheck-passed
|
||||
- check-success=full-suite-passed
|
||||
|
||||
pull_request_rules:
|
||||
|
||||
@@ -272,24 +274,15 @@ pull_request_rules:
|
||||
merge:
|
||||
method: squash
|
||||
|
||||
- name: auto-rebase when ready and Full Suite passed
|
||||
- name: auto-update when ready
|
||||
conditions:
|
||||
- label=ready
|
||||
- "#approved-reviews-by>=1"
|
||||
- check-success=full-suite-passed
|
||||
- -conflict
|
||||
- -closed
|
||||
- -draft
|
||||
actions:
|
||||
rebase: {}
|
||||
|
||||
- name: remove ready label on Full Suite failure
|
||||
conditions:
|
||||
- label=ready
|
||||
- check-failure=full-suite-passed
|
||||
actions:
|
||||
label:
|
||||
remove: [ready]
|
||||
update: {}
|
||||
|
||||
# ============================================================
|
||||
# PR title format help
|
||||
@@ -319,3 +312,5 @@ pull_request_rules:
|
||||
|
||||
Please update your PR title and the merge protection check will pass automatically.
|
||||
|
||||
merge_protections_settings:
|
||||
reporting_method: check-runs
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
name: Aggregate Test Status
|
||||
|
||||
on:
|
||||
status:
|
||||
|
||||
permissions:
|
||||
statuses: write
|
||||
|
||||
jobs:
|
||||
aggregate:
|
||||
if: >-
|
||||
github.event.context == 'direct-test-completed'
|
||||
&& github.event.state == 'success'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check and update aggregate status
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const sha = context.payload.sha;
|
||||
|
||||
const { data } = await github.rest.repos.getCombinedStatusForRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: sha,
|
||||
per_page: 100,
|
||||
});
|
||||
|
||||
const bkStatuses = data.statuses.filter(
|
||||
s => s.context.startsWith('buildkite/ci/')
|
||||
);
|
||||
|
||||
const FASTCHECK_PREFIX = 'buildkite/ci/microscope-';
|
||||
const FULL_SUITE_PREFIXES = [
|
||||
'buildkite/ci/test-tube-',
|
||||
'buildkite/ci/bar-chart-',
|
||||
];
|
||||
|
||||
const fastcheck = bkStatuses.filter(
|
||||
s => s.context.startsWith(FASTCHECK_PREFIX)
|
||||
);
|
||||
const fullSuite = bkStatuses.filter(
|
||||
s => FULL_SUITE_PREFIXES.some(p => s.context.startsWith(p))
|
||||
);
|
||||
|
||||
if (
|
||||
fastcheck.length > 0
|
||||
&& fastcheck.every(s => s.state === 'success')
|
||||
) {
|
||||
core.info(
|
||||
`All ${fastcheck.length} fastcheck tests passed — updating fastcheck-passed`
|
||||
);
|
||||
await github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha,
|
||||
state: 'success',
|
||||
context: 'fastcheck-passed',
|
||||
description:
|
||||
`All ${fastcheck.length} fastcheck tests passed`,
|
||||
});
|
||||
}
|
||||
|
||||
if (
|
||||
fullSuite.length > 0
|
||||
&& fullSuite.every(s => s.state === 'success')
|
||||
) {
|
||||
core.info(
|
||||
`All ${fullSuite.length} full suite tests passed — updating full-suite-passed`
|
||||
);
|
||||
await github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha,
|
||||
state: 'success',
|
||||
context: 'full-suite-passed',
|
||||
description:
|
||||
`All ${fullSuite.length} full suite tests passed`,
|
||||
});
|
||||
}
|
||||
@@ -4,10 +4,11 @@ on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
|
||||
concurrency:
|
||||
group: pre-commit-${{ github.ref }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
|
||||
inputs:
|
||||
ref:
|
||||
description: 'Git ref to checkout (defaults to github.ref)'
|
||||
required: false
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -18,6 +19,8 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -7,6 +7,7 @@ on:
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
statuses: write
|
||||
|
||||
jobs:
|
||||
handle-merge:
|
||||
@@ -32,6 +33,7 @@ jobs:
|
||||
core.setOutput('has_write', String(hasWrite));
|
||||
|
||||
- name: Add ready label and react
|
||||
id: label
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
@@ -39,7 +41,6 @@ jobs:
|
||||
const owner = context.repo.owner;
|
||||
const repo = context.repo.repo;
|
||||
const prNumber = context.payload.issue.number;
|
||||
// Remove ready first to allow re-trigger (labeled event fires on add, not if already present)
|
||||
try { await github.rest.issues.removeLabel({ owner, repo, issue_number: prNumber, name: 'ready' }); } catch {}
|
||||
await github.rest.issues.addLabels({ owner, repo, issue_number: prNumber, labels: ['ready'] });
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
@@ -47,6 +48,44 @@ jobs:
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
const { data: pr } = await github.rest.pulls.get({ owner, repo, pull_number: prNumber });
|
||||
core.setOutput('pr_sha', pr.head.sha);
|
||||
core.setOutput('pr_branch', pr.head.ref);
|
||||
core.setOutput('pr_number', String(prNumber));
|
||||
|
||||
- name: Trigger Full Suite
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_SHA: ${{ steps.label.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
curl -sS --fail-with-body -X POST \
|
||||
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds" \
|
||||
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
-H "Content-Type: application/json" \
|
||||
--data-raw "$(jq -n \
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
branch: $branch,
|
||||
message: $message,
|
||||
ignore_pipeline_branch_filters: true,
|
||||
pull_request_id: $pr_id,
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
github.event.issue.pull_request != null
|
||||
@@ -142,21 +181,8 @@ jobs:
|
||||
core.setOutput('sha', pr.head.sha);
|
||||
core.setOutput('branch', pr.head.ref);
|
||||
|
||||
pre-commit:
|
||||
needs: parse-command
|
||||
if: >-
|
||||
needs.parse-command.outputs.has_write == 'true'
|
||||
&& needs.parse-command.outputs.test_scope == 'precommit'
|
||||
uses: ./.github/workflows/ci-precommit.yml
|
||||
|
||||
trigger-buildkite:
|
||||
needs: parse-command
|
||||
if: >-
|
||||
needs.parse-command.outputs.has_write == 'true'
|
||||
&& needs.parse-command.outputs.test_type != ''
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: React to comment
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
@@ -167,6 +193,43 @@ jobs:
|
||||
content: 'rocket',
|
||||
});
|
||||
|
||||
pre-commit:
|
||||
needs: parse-command
|
||||
if: >-
|
||||
needs.parse-command.outputs.has_write == 'true'
|
||||
&& needs.parse-command.outputs.test_scope == 'precommit'
|
||||
uses: ./.github/workflows/ci-precommit.yml
|
||||
with:
|
||||
ref: refs/pull/${{ github.event.issue.number }}/merge
|
||||
|
||||
post-precommit-status:
|
||||
needs: [parse-command, pre-commit]
|
||||
if: always() && needs.parse-command.outputs.test_scope == 'precommit'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
env:
|
||||
PR_SHA: ${{ needs.parse-command.outputs.pr_sha }}
|
||||
RESULT: ${{ needs.pre-commit.result }}
|
||||
with:
|
||||
script: |
|
||||
const state = process.env.RESULT === 'success' ? 'success' : 'failure';
|
||||
await github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha: process.env.PR_SHA,
|
||||
state,
|
||||
context: 'pre-commit',
|
||||
description: `Triggered via /test pre-commit (${state})`,
|
||||
});
|
||||
|
||||
trigger-buildkite:
|
||||
needs: parse-command
|
||||
if: >-
|
||||
needs.parse-command.outputs.has_write == 'true'
|
||||
&& needs.parse-command.outputs.test_type != ''
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Trigger Buildkite
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
name: Trigger Full Suite
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
pull_request_target:
|
||||
types: [labeled, synchronize]
|
||||
|
||||
permissions:
|
||||
@@ -10,7 +10,7 @@ permissions:
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
trigger:
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
# Find running builds for this branch with TEST_SCOPE=full and cancel them
|
||||
builds=$(curl -sS -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds?branch=${PR_BRANCH}&state=running,scheduled" \
|
||||
| jq -r '.[] | select(.env.TEST_SCOPE == "full") | .number')
|
||||
| jq -r '.[] | select(try (.env.TEST_SCOPE == "full") catch false) | .number')
|
||||
for build_num in $builds; do
|
||||
echo "Cancelling Buildkite build #$build_num"
|
||||
curl -sS -X PUT -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
|
||||
@@ -85,6 +85,7 @@ docs/distillation/examples/
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
# Next.js / Node artifacts under ui/: see ui/.gitignore
|
||||
|
||||
.claude/
|
||||
.codex/
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
WRN 2026-03-26T13:46:33.469 ?.19646 server_start:193: Failed to start server: operation not permitted: /var/folders/z_/h_6myyk14d1b7z87z3vy4mjh0000gn/T/nvim.dsynkd/iSe0el/nvim.19646.0
|
||||
@@ -0,0 +1 @@
|
||||
3.12
|
||||
@@ -24,7 +24,7 @@ PR push
|
||||
Runs on the PR branch directly
|
||||
│
|
||||
pass ──► Mergify auto-squash-merges to main, branch deleted
|
||||
fail ──► Mergify removes 'ready' label; fix and /merge again
|
||||
fail ──► fix the regression, push, and /merge again
|
||||
```
|
||||
|
||||
---
|
||||
@@ -102,8 +102,8 @@ failing test's output.
|
||||
| Performance Tests | `performance` | 30 min |
|
||||
| API Server Tests | `api_server` | 30 min |
|
||||
|
||||
A Full Suite failure removes the `ready` label automatically. A Mergify comment links to
|
||||
the Buildkite build. Fix the regression, push, and comment `/merge` again.
|
||||
If a Full Suite test fails, check the Buildkite build log for the failing step's output.
|
||||
Fix the regression, push, and comment `/merge` again to re-trigger.
|
||||
|
||||
---
|
||||
|
||||
@@ -129,8 +129,8 @@ Suite passing directly on the PR branch.
|
||||
- No merge conflicts
|
||||
5. If all conditions pass, Mergify squash-merges to `main` automatically. The branch is
|
||||
deleted after merge.
|
||||
6. If the Full Suite fails, Mergify removes the `ready` label and posts a comment linking to
|
||||
the Buildkite build. The developer fixes the issue, pushes, and comments `/merge` again.
|
||||
6. If the Full Suite fails, the developer fixes the issue, pushes, and comments `/merge`
|
||||
again to re-trigger.
|
||||
|
||||
**Merge conditions summary:**
|
||||
|
||||
@@ -279,6 +279,30 @@ Triggers a specific Buildkite test or suite on the current PR branch.
|
||||
| `/test api` | API server integration tests | `api_server` |
|
||||
| `/test full` | Entire Full Suite | all (with `TEST_SCOPE=full`) |
|
||||
| `/test fastcheck` | Entire Fastcheck suite | fastcheck (with `TEST_SCOPE=fastcheck`) |
|
||||
| `/test pre-commit` | Pre-commit checks on PR code | — (runs `ci-precommit.yml` via `workflow_call`) |
|
||||
|
||||
**Re-running failed tests:** When you use `/test <name>` to re-run a specific failed test,
|
||||
the resulting Buildkite check uses the same name as the original (e.g., `/test encoder`
|
||||
creates `buildkite/ci/microscope-encoder-tests`). This overwrites the failed check status.
|
||||
Once all tests in a tier pass, the aggregate status (`fastcheck-passed` or
|
||||
`full-suite-passed`) is automatically updated to `success` by the `ci-aggregate-status.yml`
|
||||
workflow.
|
||||
|
||||
**How aggregate status refresh works:**
|
||||
|
||||
1. `/test <name>` triggers a Buildkite build with `TEST_SCOPE=direct`. The test step uses
|
||||
the same label as its fastcheck/full-suite counterpart, so the resulting GitHub check
|
||||
overwrites the original.
|
||||
2. When the build completes, Buildkite's `notify` posts a `direct-test-completed` commit
|
||||
status. This is the only signal that triggers the aggregate workflow — intermediate step
|
||||
status updates do not trigger it.
|
||||
3. `ci-aggregate-status.yml` fires, calls `getCombinedStatusForRef` to fetch the latest
|
||||
status for every context on that commit (each context returns only its most recent
|
||||
state), groups them by prefix (`microscope-*` → fastcheck, `test-tube-*`/`bar-chart-*`
|
||||
→ full suite), and posts `fastcheck-passed: success` or `full-suite-passed: success` if
|
||||
all entries in the group are `success`.
|
||||
4. Tests that were never triggered (skipped by monorepo-diff) have no status entry and do
|
||||
not block the aggregate.
|
||||
|
||||
---
|
||||
|
||||
@@ -296,6 +320,7 @@ Protected branches (`main`, `master`, `release/*`) are never deleted.
|
||||
| `ci-precommit.yml` | Every push / PR against `main` | Runs pre-commit hooks (yapf, ruff, mypy, codespell, pymarkdown, actionlint, check-filenames) |
|
||||
| `ci-trigger-full-suite.yml` | `ready` label added to a PR | Calls Buildkite API to run Full Suite on the PR branch |
|
||||
| `ci-slash-commands.yml` | PR comment starting with `/merge` or `/test` | Handles slash commands; adds `ready` label or triggers Buildkite |
|
||||
| `ci-aggregate-status.yml` | Any Buildkite commit status update | Checks if all tests in a tier passed; updates `fastcheck-passed` or `full-suite-passed` |
|
||||
| `community-issue-labeler.yml` | Issue opened or edited | Auto-labels issues by keyword matching against title and body |
|
||||
| `community-welcome.yml` | First contribution | Posts a welcome comment for first-time contributors |
|
||||
| `community-stale.yml` | Scheduled | Marks and closes stale issues and PRs |
|
||||
|
||||
@@ -104,8 +104,9 @@ distillation, self-forcing, VSA, VMoBA, performance benchmarks, and API server t
|
||||
8. If all Full Suite tests pass and all merge conditions are met (approval, valid title,
|
||||
pre-commit green, fastcheck green, no draft, no conflicts), Mergify squash-merges to
|
||||
`main` automatically. Your branch is deleted.
|
||||
9. If a Full Suite test fails, Mergify removes the `ready` label and posts a comment with a
|
||||
link to the Buildkite build. Fix the issue, push, and comment `/merge` again.
|
||||
9. If a Full Suite test fails, check the Buildkite build log for the failing step. Fix the
|
||||
issue, push, and comment `/merge` again. You can also re-run individual failed tests
|
||||
with `/test <name>` — see below.
|
||||
|
||||
!!! note
|
||||
Only contributors with write permission to the repository can trigger slash commands.
|
||||
@@ -149,10 +150,15 @@ Comment on your PR to trigger specific tests independently of the auto-merge flo
|
||||
/test vmoba # VMoBA inference tests
|
||||
/test performance # Performance benchmarks
|
||||
/test api # API server integration tests
|
||||
/test pre-commit # Pre-commit checks on PR code
|
||||
```
|
||||
|
||||
The workflow reacts with a 🚀 emoji to confirm the command was received.
|
||||
|
||||
When you re-run an individual test with `/test <name>`, the new result overwrites the
|
||||
original failed check (same Buildkite check name). Once all tests in a tier pass, the
|
||||
`fastcheck-passed` or `full-suite-passed` status is automatically updated.
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
@@ -199,9 +205,8 @@ Mergify removes the `needs-rebase` label automatically once conflicts are resolv
|
||||
|
||||
### Full Suite failed after `/merge`
|
||||
|
||||
The Full Suite found a regression. Mergify removes the `ready` label and posts a comment
|
||||
linking to the Buildkite build. Check the failing step's output for assertion errors or
|
||||
tracebacks.
|
||||
The Full Suite found a regression. Check the failing Buildkite step's output for assertion
|
||||
errors or tracebacks.
|
||||
|
||||
Common causes:
|
||||
|
||||
|
||||
@@ -0,0 +1,508 @@
|
||||
status_definitions:
|
||||
kept: "Public field remains on a public adapter surface with the same meaning."
|
||||
moved: "Public field remains supported but normalizes into a different nested path."
|
||||
profile_owned: "Public field remains supported only through a model/profile-specific surface."
|
||||
compatibility_only: "Legacy public field remains adapter-only during migration and is not part of the canonical typed schema."
|
||||
private_only: "Field should only be handled by private adapters and is not a public FastVideo compatibility promise."
|
||||
internal_only: "Field is runtime/config plumbing and should not be part of the new public typed inference API."
|
||||
|
||||
surfaces:
|
||||
fastvideo_args:
|
||||
moved:
|
||||
model_path: generator.model_path
|
||||
workload_type: generator.pipeline.workload_type
|
||||
distributed_executor_backend: generator.engine.execution_backend
|
||||
trust_remote_code: generator.trust_remote_code
|
||||
revision: generator.revision
|
||||
num_gpus: generator.engine.num_gpus
|
||||
tp_size: generator.engine.parallelism.tp_size
|
||||
sp_size: generator.engine.parallelism.sp_size
|
||||
hsdp_replicate_dim: generator.engine.parallelism.hsdp_replicate_dim
|
||||
hsdp_shard_dim: generator.engine.parallelism.hsdp_shard_dim
|
||||
dist_timeout: generator.engine.parallelism.dist_timeout
|
||||
lora_path: generator.pipeline.components.lora_path
|
||||
dit_cpu_offload: generator.engine.offload.dit
|
||||
use_fsdp_inference: generator.engine.use_fsdp_inference
|
||||
dit_layerwise_offload: generator.engine.offload.dit_layerwise
|
||||
text_encoder_cpu_offload: generator.engine.offload.text_encoder
|
||||
image_encoder_cpu_offload: generator.engine.offload.image_encoder
|
||||
vae_cpu_offload: generator.engine.offload.vae
|
||||
pin_cpu_memory: generator.engine.offload.pin_cpu_memory
|
||||
enable_torch_compile: generator.engine.compile.enabled
|
||||
torch_compile_kwargs: generator.engine.compile.kwargs
|
||||
disable_autocast: generator.engine.disable_autocast
|
||||
enable_stage_verification: generator.engine.enable_stage_verification
|
||||
prompt_txt: request.inputs.prompt_path
|
||||
override_text_encoder_safetensors: generator.pipeline.components.text_encoder_weights
|
||||
override_text_encoder_quant: generator.engine.quantization.text_encoder_quant
|
||||
override_transformer_cls_name: generator.pipeline.components.override_transformer_cls_name
|
||||
init_weights_from_safetensors: generator.pipeline.components.transformer_weights
|
||||
init_weights_from_safetensors_2: generator.pipeline.components.transformer_2_weights
|
||||
override_pipeline_cls_name: generator.pipeline.components.override_pipeline_cls_name
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
profile_owned:
|
||||
ltx2_vae_tiling: generator.pipeline.profile_overrides.ltx2.vae_tiling
|
||||
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_size_in_pixels
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
|
||||
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_size_in_frames
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_overlap_in_frames
|
||||
ltx2_initial_latent_path: request.extensions.ltx2.initial_latent_path
|
||||
compatibility_only:
|
||||
mode: "Legacy multi-mode FastVideoArgs switch; typed inference config should not expose execution mode."
|
||||
inference_mode: "Legacy boolean mirror of mode; kept only through adapters while FastVideoArgs remains."
|
||||
lora_nickname: "Legacy adapter-selection surface pending LoRA API cleanup."
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
private_only:
|
||||
ray_placement_group: "Ray deployment-only field."
|
||||
ray_runtime_env: "Ray deployment-only field."
|
||||
internal_only:
|
||||
pipeline_config: "Legacy internal carrier object."
|
||||
preprocess_config: "Legacy preprocess carrier object."
|
||||
moba_config: "Derived runtime config loaded from moba_config_path."
|
||||
model_paths: "Runtime bookkeeping."
|
||||
model_loaded: "Runtime bookkeeping."
|
||||
|
||||
pipeline_config_base:
|
||||
moved:
|
||||
pipeline_config_path: generator.pipeline.components.pipeline_config_path
|
||||
profile_owned:
|
||||
embedded_cfg_scale: generator.pipeline.profile_overrides.embedded_cfg_scale
|
||||
flow_shift: generator.pipeline.profile_overrides.flow_shift
|
||||
flow_shift_sr: generator.pipeline.profile_overrides.flow_shift_sr
|
||||
is_causal: generator.pipeline.profile_overrides.is_causal
|
||||
vae_tiling: generator.pipeline.profile_overrides.vae_tiling
|
||||
vae_sp: generator.pipeline.profile_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.profile_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.profile_overrides.ti2v_task
|
||||
boundary_ratio: generator.pipeline.profile_overrides.boundary_ratio
|
||||
compatibility_only:
|
||||
model_path: "Redundant with generator.model_path."
|
||||
disable_autocast: "Duplicated by generator.engine.disable_autocast during migration."
|
||||
dit_precision: "Precision override pending dedicated typed component precision design."
|
||||
upsampler_precision: "Precision override pending dedicated typed component precision design."
|
||||
vae_precision: "Precision override pending dedicated typed component precision design."
|
||||
image_encoder_precision: "Precision override pending dedicated typed component precision design."
|
||||
text_encoder_precisions: "Precision override pending dedicated typed component precision design."
|
||||
internal_only:
|
||||
dit_config: "Legacy internal component config object."
|
||||
upsampler_config: "Legacy internal component config object."
|
||||
vae_config: "Legacy internal component config object."
|
||||
image_encoder_config: "Legacy internal component config object."
|
||||
text_encoder_configs: "Legacy internal component config object."
|
||||
preprocess_text_funcs: "Internal text preprocessing hooks."
|
||||
postprocess_text_funcs: "Internal text postprocessing hooks."
|
||||
|
||||
pipeline_config_extensions:
|
||||
profile_owned:
|
||||
conditioning_strategy:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
max_num_conditional_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
min_num_conditional_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
sigma_conditional:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
sigma_data:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
state_ch:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
state_t:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
text_encoder_class:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
autoregressive_chunk_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
autoregressive_overlap_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
cfg_behavior:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_camera_rotation:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_movement_distance:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_negative_prompt:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_trajectory_type:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
filter_points_threshold:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
fps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
frame_buffer_max:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
moge_model_name:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
noise_aug_strength:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
num_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
offload_moge_after_depth:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
use_moge_depth:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
video_resolution:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
text_encoder_crop_start:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V480PStepDistilledConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V720PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15SR1080PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V480PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V720PConfig
|
||||
- fastvideo.configs.pipelines.hyworld.HYWorldConfig
|
||||
- fastvideo.configs.pipelines.hyworld.Hunyuan15T2V480PConfig
|
||||
text_encoder_max_lengths:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V480PStepDistilledConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V720PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15SR1080PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V480PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V720PConfig
|
||||
- fastvideo.configs.pipelines.hyworld.HYWorldConfig
|
||||
- fastvideo.configs.pipelines.hyworld.Hunyuan15T2V480PConfig
|
||||
precision:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.lingbotworld.LingBotWorldI2V480PConfig
|
||||
- fastvideo.configs.pipelines.lingbotworld.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_1_3B_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_1_T2V_480P_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameBaseI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWan2_2_T2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WANV2VConfig
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_T2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.WanI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanI2V720PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V720PConfig
|
||||
warp_denoising_step:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.lingbotworld.LingBotWorldI2V480PConfig
|
||||
- fastvideo.configs.pipelines.lingbotworld.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_1_3B_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_1_T2V_480P_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameBaseI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWan2_2_T2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WANV2VConfig
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_T2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.WanI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanI2V720PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V720PConfig
|
||||
bsa_cdf_threshold:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_chunk_k:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_chunk_q:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_params:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_sparsity:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enable_bsa:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enable_kv_cache:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enhance_hf:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
offload_kv_cache:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
t_thresh:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
use_distill:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
scheduler_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
text_encoder_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
tokenizer_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
transformer_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
vae_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
expand_timesteps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
context_noise:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
num_frames_per_block:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
compatibility_only:
|
||||
batch_size: "Gen3C inference-only tuning field pending typed batching design."
|
||||
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
|
||||
guidance_scale: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
num_inference_steps: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
internal_only:
|
||||
audio_decoder_config: "Legacy internal component config object."
|
||||
audio_decoder_precision: "Precision override pending dedicated component precision design."
|
||||
vocoder_config: "Legacy internal component config object."
|
||||
vocoder_precision: "Precision override pending dedicated component precision design."
|
||||
|
||||
sampling_param_base:
|
||||
moved:
|
||||
image_path: request.inputs.image_path
|
||||
pil_image: request.inputs.pil_image
|
||||
video_path: request.inputs.video_path
|
||||
mouse_cond: request.inputs.mouse_cond
|
||||
keyboard_cond: request.inputs.keyboard_cond
|
||||
grid_sizes: request.inputs.grid_sizes
|
||||
pose: request.inputs.pose
|
||||
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
|
||||
refine_from: request.inputs.refine_from
|
||||
stage1_video: request.inputs.stage1_video
|
||||
prompt: request.prompt
|
||||
negative_prompt: request.negative_prompt
|
||||
prompt_path: request.inputs.prompt_path
|
||||
output_path: request.output.output_path
|
||||
output_video_name: request.output.output_video_name
|
||||
num_videos_per_prompt: request.sampling.num_videos_per_prompt
|
||||
seed: request.sampling.seed
|
||||
num_frames: request.sampling.num_frames
|
||||
height: request.sampling.height
|
||||
width: request.sampling.width
|
||||
height_sr: request.sampling.height_sr
|
||||
width_sr: request.sampling.width_sr
|
||||
fps: request.sampling.fps
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
num_inference_steps_sr: request.sampling.num_inference_steps_sr
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
save_video: request.output.save_video
|
||||
return_frames: request.output.return_frames
|
||||
return_trajectory_latents: request.runtime.return_trajectory_latents
|
||||
return_trajectory_decoded: request.runtime.return_trajectory_decoded
|
||||
profile_owned:
|
||||
t_thresh: request.stage_overrides.refine.t_thresh
|
||||
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
|
||||
num_cond_frames: request.stage_overrides.refine.num_cond_frames
|
||||
trajectory_type: request.extensions.gen3c.trajectory_type
|
||||
movement_distance: request.extensions.gen3c.movement_distance
|
||||
camera_rotation: request.extensions.gen3c.camera_rotation
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
sampling_param_extensions:
|
||||
moved:
|
||||
guidance_scale_2:
|
||||
target: request.sampling.guidance_scale_2
|
||||
sources:
|
||||
- fastvideo.configs.sample.lingbotworld.LingBotWorld_SamplingParam
|
||||
- fastvideo.configs.sample.lingbotworld.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.SelfForcingWan2_2_T2V_A14B_480P_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_T2V_A14B_SamplingParam
|
||||
profile_owned:
|
||||
action_list:
|
||||
target: request.extensions.hunyuangamecraft.action_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
action_speed_list:
|
||||
target: request.extensions.hunyuangamecraft.action_speed_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_states:
|
||||
target: request.extensions.hunyuangamecraft.camera_states
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_trajectory:
|
||||
target: request.extensions.hunyuangamecraft.camera_trajectory
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
conditioning_mask:
|
||||
target: request.extensions.hunyuangamecraft.conditioning_mask
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
gt_latents:
|
||||
target: request.extensions.hunyuangamecraft.gt_latents
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
prompt_attention_mask:
|
||||
target: request.extensions.hyworld.prompt_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
negative_attention_mask:
|
||||
target: request.extensions.hyworld.negative_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
ltx2_cfg_scale_audio:
|
||||
target: request.extensions.ltx2.cfg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_cfg_scale_video:
|
||||
target: request.extensions.ltx2.cfg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_audio:
|
||||
target: request.extensions.ltx2.modality_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_video:
|
||||
target: request.extensions.ltx2.modality_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_rescale_scale:
|
||||
target: request.extensions.ltx2.rescale_scale
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_audio:
|
||||
target: request.extensions.ltx2.stg_blocks_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_video:
|
||||
target: request.extensions.ltx2.stg_blocks_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_audio:
|
||||
target: request.extensions.ltx2.stg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_video:
|
||||
target: request.extensions.ltx2.stg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
|
||||
openai_image_request:
|
||||
kept:
|
||||
model: "HTTP adapter model-routing field."
|
||||
response_format: "HTTP adapter response formatting field."
|
||||
output_format: "HTTP adapter output-format field."
|
||||
background: "HTTP adapter output-format field."
|
||||
quality: "Compatibility field currently accepted by the adapter."
|
||||
style: "Compatibility field currently accepted by the adapter."
|
||||
user: "Compatibility field currently accepted by the adapter."
|
||||
moved:
|
||||
prompt: request.prompt
|
||||
n: request.sampling.num_videos_per_prompt
|
||||
size:
|
||||
target: request.sampling.width,height
|
||||
note: "Adapter parses OpenAI size strings as WIDTHxHEIGHT and forwards width then height."
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
seed: request.sampling.seed
|
||||
negative_prompt: request.negative_prompt
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
openai_video_request:
|
||||
kept:
|
||||
model: "HTTP adapter model-routing field."
|
||||
moved:
|
||||
prompt: request.prompt
|
||||
input_reference: request.inputs.image_path
|
||||
reference_url: request.inputs.image_path
|
||||
size:
|
||||
target: request.sampling.width,height
|
||||
note: "Adapter parses OpenAI size strings as WIDTHxHEIGHT and forwards width then height."
|
||||
fps: request.sampling.fps
|
||||
num_frames: request.sampling.num_frames
|
||||
seed: request.sampling.seed
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
negative_prompt: request.negative_prompt
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
output_path: request.output.output_path
|
||||
compatibility_only:
|
||||
seconds:
|
||||
target: request.sampling.num_frames
|
||||
note: "HTTP adapter duration convenience field. If num_frames is omitted, the adapter computes num_frames = fps * seconds."
|
||||
|
||||
cli:
|
||||
notes:
|
||||
- "CLI parity is checked against the actual generate/serve parser dest sets."
|
||||
- "The inventory tracks parser dest names, excluding argparse's implicit help action."
|
||||
- "The refactored inference CLI is config-only: subcommands expose only --config, and any additional CLI input must use dotted override paths."
|
||||
generate:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
expected_dests:
|
||||
- config
|
||||
serve:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
expected_dests:
|
||||
- config
|
||||
@@ -16,7 +16,8 @@ Both models are trained on **61×448×832** resolution but support generating vi
|
||||
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN \
|
||||
fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml
|
||||
```
|
||||
|
||||
## 🗂️ Dataset
|
||||
|
||||
@@ -455,6 +455,6 @@ User: generator.generate_video(prompt, ...)
|
||||
`fastvideo/pipelines/stages/`, implement `forward()`, optionally
|
||||
implement `verify_input()`/`verify_output()`.
|
||||
|
||||
7. **Verify** — Run `fastvideo generate --model-path <path> --prompt
|
||||
"test" --num-inference-steps 2` to confirm the pipeline loads and
|
||||
generates output.
|
||||
7. **Verify** — Run `fastvideo generate --config <config.yaml>` with a
|
||||
minimal nested config to confirm the pipeline loads and generates
|
||||
output.
|
||||
|
||||
+42
-81
@@ -1,71 +1,29 @@
|
||||
# FastVideo CLI Inference
|
||||
|
||||
The FastVideo CLI exposes the same core inference controls as the Python API.
|
||||
The FastVideo CLI is config-first. Inference runs are driven by a nested JSON or
|
||||
YAML config, with optional dotted-path overrides on the command line. The
|
||||
contract matches training: use an explicit subcommand plus `--config`, then add
|
||||
any dotted overrides you need.
|
||||
|
||||
## Basic Usage
|
||||
|
||||
Use either:
|
||||
|
||||
1. `--model-path` + `--prompt`
|
||||
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
|
||||
3. `--config` (JSON/YAML)
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt "A cat playing with a ball of yarn"
|
||||
fastvideo generate --config config.yaml
|
||||
fastvideo serve --config serve.yaml
|
||||
```
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt-txt prompts.txt
|
||||
```
|
||||
|
||||
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
|
||||
|
||||
## View All Arguments
|
||||
|
||||
```bash
|
||||
fastvideo generate --help
|
||||
```
|
||||
|
||||
Arguments come from:
|
||||
The subcommands intentionally expose only `--config`. Any per-run CLI changes
|
||||
must use dotted override paths such as:
|
||||
|
||||
- FastVideo runtime args (`FastVideoArgs`)
|
||||
- Sampling args (`SamplingParam`)
|
||||
- Pipeline config args (`PipelineConfig`)
|
||||
|
||||
## Common Arguments
|
||||
|
||||
### Parallelism
|
||||
|
||||
- `--num-gpus`
|
||||
- `--sp-size`
|
||||
- `--tp-size`
|
||||
|
||||
### Sampling
|
||||
|
||||
- `--num-frames`
|
||||
- `--height` / `--width`
|
||||
- `--num-inference-steps`
|
||||
- `--guidance-scale`
|
||||
- `--seed`
|
||||
- `--negative-prompt`
|
||||
|
||||
### Output
|
||||
|
||||
- `--output-path`
|
||||
- `--save-video` / `--no-save-video`
|
||||
- `--return-frames`
|
||||
|
||||
### Offloading and Performance
|
||||
|
||||
- `--dit-layerwise-offload`
|
||||
- `--use-fsdp-inference`
|
||||
- `--text-encoder-cpu-offload`
|
||||
- `--image-encoder-cpu-offload`
|
||||
- `--vae-cpu-offload`
|
||||
- `--enable-torch-compile`
|
||||
- `--torch-compile-kwargs`
|
||||
- `--generator.engine.num_gpus 2`
|
||||
- `--request.sampling.seed 42`
|
||||
- `--server.port 9000`
|
||||
|
||||
## Using Config Files
|
||||
|
||||
@@ -73,50 +31,53 @@ Arguments come from:
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config files can be JSON or YAML. CLI flags override config-file values.
|
||||
Config files can be JSON or YAML. Dotted CLI overrides take precedence over
|
||||
config-file values.
|
||||
|
||||
Example `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara lounging in a hammock"
|
||||
output_path: "outputs/"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
request:
|
||||
prompt: A capybara lounging in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- Use `dit_precision` / `vae_precision` (not `precision`).
|
||||
- Nested config objects are supported, for example `vae_config` and
|
||||
`dit_config`.
|
||||
- `generator` and `request` are the top-level keys for generation configs.
|
||||
- `serve` configs use `generator`, `server`, and optional `default_request`.
|
||||
- Prompt text files belong under `request.inputs.prompt_path`.
|
||||
|
||||
## Examples
|
||||
|
||||
Simple generation:
|
||||
|
||||
```bash
|
||||
fastvideo generate \
|
||||
--model-path FastVideo/FastHunyuan-diffusers \
|
||||
--prompt "A cat playing with a ball of yarn" \
|
||||
--num-frames 45 --height 720 --width 1280 \
|
||||
--num-inference-steps 6 --seed 1024 \
|
||||
--output-path outputs/
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config + CLI override:
|
||||
Config + dotted override:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
|
||||
fastvideo generate --config config.yaml --request.prompt "A panda skiing at sunset"
|
||||
```
|
||||
|
||||
Helper wrapper with positional config path:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/run.sh scripts/inference/inference_wan.yaml
|
||||
```
|
||||
|
||||
@@ -73,32 +73,40 @@ if __name__ == '__main__':
|
||||
|
||||
## JSON/YAML Config Files (CLI)
|
||||
|
||||
The CLI supports `--config` with JSON or YAML. Command-line arguments override
|
||||
config file values.
|
||||
By default, `fastvideo generate` uses `return_frames=false` unless you set
|
||||
`--return-frames` (or `return_frames: true` in config).
|
||||
The inference CLI is config-first. Use an explicit subcommand with `--config`,
|
||||
then apply optional dotted overrides on top, matching the training CLI style.
|
||||
By default, CLI generation uses `return_frames=false` unless you set
|
||||
`request.output.return_frames: true` in config or via a dotted override.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Use CLI argument names as keys (underscore or hyphen is accepted). Example:
|
||||
Example nested config:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara relaxing in a hammock"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
request:
|
||||
prompt: A capybara relaxing in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Override individual values from the CLI with dotted paths:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --request.sampling.seed 42
|
||||
```
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
# GEN3C: 3D-Informed Camera-Controlled Video Generation
|
||||
|
||||
[GEN3C](https://arxiv.org/abs/2503.03751) is NVIDIA's Cosmos-7B-based video model for camera-controlled generation from a single image. The FastVideo integration supports the GEN3C I2V workflow, including 3D cache conditioning and tokenizer-based conditioning latents.
|
||||
|
||||
## Key Features
|
||||
|
||||
- **Camera trajectory control**: `left/right/up/down/zoom_in/zoom_out/clockwise/counterclockwise`
|
||||
- **3D cache conditioning**: depth prediction -> point cloud cache -> forward warping -> latent conditioning
|
||||
- **Single-image to video generation**: 121-frame generation with camera motion
|
||||
- **Official raw checkpoint conversion**: `model.pt` -> Diffusers/FastVideo layout
|
||||
|
||||
## Model Sources
|
||||
|
||||
- Official raw checkpoint (not Diffusers): `nvidia/GEN3C-Cosmos-7B`
|
||||
- Diffusers-format checkpoint: `FastVideo/GEN3C-Cosmos-7B-Diffusers`
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
```
|
||||
|
||||
- If you hit `ImportError: libGL.so.1` (common on Ubuntu/headless nodes), you can try installing OpenCV runtime libs:
|
||||
|
||||
```bash
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Option A: Use Diffusers-format weights directly
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/basic_gen3c.py \
|
||||
--model_path FastVideo/GEN3C-Cosmos-7B-Diffusers \
|
||||
--image_path /path/to/input.png \
|
||||
--prompt "" \
|
||||
--trajectory left \
|
||||
--movement_distance 0.3 \
|
||||
--camera_rotation center_facing \
|
||||
--num_inference_steps 35 \
|
||||
--guidance_scale 1.0 \
|
||||
--output_path outputs_video/gen3c_output.mp4
|
||||
```
|
||||
|
||||
### Option B: Convert official raw checkpoint locally
|
||||
|
||||
1. Download:
|
||||
|
||||
```bash
|
||||
huggingface-cli download nvidia/GEN3C-Cosmos-7B --local-dir official_weights/GEN3C-Cosmos-7B
|
||||
```
|
||||
|
||||
1. Convert:
|
||||
|
||||
```bash
|
||||
python scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py \
|
||||
--source official_weights/GEN3C-Cosmos-7B/model.pt \
|
||||
--output converted_weights/GEN3C-Cosmos-7B
|
||||
```
|
||||
|
||||
1. Run:
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/basic_gen3c.py \
|
||||
--model_path converted_weights/GEN3C-Cosmos-7B \
|
||||
--image_path /path/to/input.png \
|
||||
--prompt "" \
|
||||
--trajectory left \
|
||||
--movement_distance 0.3 \
|
||||
--camera_rotation center_facing \
|
||||
--num_inference_steps 35 \
|
||||
--guidance_scale 1.0 \
|
||||
--output_path outputs_video/gen3c_output.mp4
|
||||
```
|
||||
|
||||
## FastVideo Defaults
|
||||
|
||||
GEN3C defaults in FastVideo:
|
||||
|
||||
- `height=704`, `width=1280`
|
||||
- `num_frames=121`
|
||||
- `num_inference_steps=35`
|
||||
- `guidance_scale=1.0`
|
||||
- `fps=24`
|
||||
|
||||
These values are defined in:
|
||||
|
||||
- `fastvideo/configs/sample/gen3c.py`
|
||||
- `fastvideo/configs/pipelines/gen3c.py`
|
||||
|
||||
and align with the official GEN3C inference defaults in:
|
||||
|
||||
- `tmp/GEN3C/cosmos_predict1/diffusion/inference/inference_utils.py`
|
||||
|
||||
## Scheduler Note
|
||||
|
||||
The converted GEN3C Diffusers layout may include a FlowMatch scheduler config, but GEN3C denoising uses EDM preconditioning behavior. FastVideo's GEN3C pipeline enforces an EDM scheduler at runtime for parity with official inference behavior.
|
||||
|
||||
Implementation path:
|
||||
|
||||
- `fastvideo/pipelines/basic/gen3c/gen3c_pipeline.py`
|
||||
|
||||
## 3D Cache Conditioning Path
|
||||
|
||||
FastVideo GEN3C conditioning stage performs:
|
||||
|
||||
1. MoGe depth estimation from input image
|
||||
2. 3D cache initialization
|
||||
3. Camera trajectory generation
|
||||
4. Forward rendering of warped frames + masks
|
||||
5. VAE/tokenizer encoding of conditioning buffers
|
||||
6. Denoising with condition mask + condition pose channels
|
||||
|
||||
Main implementation:
|
||||
|
||||
- `fastvideo/pipelines/basic/gen3c/gen3c_pipeline.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/cache_3d.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/depth_estimation.py`
|
||||
- `fastvideo/models/vaes/gen3c_tokenizer_vae.py`
|
||||
|
||||
## References
|
||||
|
||||
- [GEN3C Paper](https://arxiv.org/abs/2503.03751)
|
||||
- [Official Repository](https://github.com/nv-tlabs/GEN3C)
|
||||
- [Official Checkpoint (raw)](https://huggingface.co/nvidia/GEN3C-Cosmos-7B)
|
||||
@@ -73,6 +73,7 @@ pipeline initialization and sampling.
|
||||
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 TempleRun | `FastVideo/Matrix-Game-2.0-TempleRun-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
@@ -85,6 +86,11 @@ The authoritative source for model-ID recognition is
|
||||
`fastvideo/registry.py`. If a model ID is registered there, FastVideo can
|
||||
resolve default pipeline and sampling configuration for it.
|
||||
|
||||
**Note (GEN3C)**: The official `nvidia/GEN3C-Cosmos-7B` repo provides a raw
|
||||
`model.pt` checkpoint. Use a Diffusers-format repo (for example,
|
||||
`FastVideo/GEN3C-Cosmos-7B-Diffusers`) or convert locally with
|
||||
`scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py`.
|
||||
|
||||
## Special requirements
|
||||
|
||||
### Sliding Tile Attention
|
||||
|
||||
@@ -28,6 +28,11 @@ For an example running DMD+VSA inference:
|
||||
python examples/inference/basic/basic_dmd.py
|
||||
```
|
||||
|
||||
For the typed config/request path added during the inference API refactor:
|
||||
```
|
||||
python examples/inference/basic/basic_dmd_new_api.py
|
||||
```
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
PipelineSelection,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2_typed"
|
||||
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
model_name = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
# PR 2 still routes a few advanced inference knobs through the
|
||||
# compatibility bridge until they get first-class typed fields.
|
||||
pipeline=PipelineSelection(
|
||||
experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic."
|
||||
)
|
||||
request2 = GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result2 = generator.generate(request2)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"First output written to: {result.video_path}")
|
||||
print(f"Time taken to generate video2: {gen_time2} seconds")
|
||||
print(f"Second output written to: {result2.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,109 @@
|
||||
"""
|
||||
GEN3C: 3D-aware camera-controlled video generation.
|
||||
|
||||
This example generates a video from a single input image with camera control.
|
||||
The pipeline uses MoGe depth estimation, 3D point cloud forward warping,
|
||||
and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
If you hit `ImportError: libGL.so.1`, install:
|
||||
sudo apt-get update && sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
2. Download and convert weights:
|
||||
huggingface-cli download nvidia/GEN3C-Cosmos-7B --local-dir official_weights/GEN3C-Cosmos-7B
|
||||
python scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py \
|
||||
--source ./official_weights/GEN3C-Cosmos-7B/model.pt \
|
||||
--output ./converted_weights/GEN3C-Cosmos-7B \
|
||||
--components-source nvidia/Cosmos-Predict2-2B-Video2World
|
||||
3. Provide an input image for 3D-conditioned generation.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="GEN3C video generation")
|
||||
parser.add_argument("--model_path",
|
||||
type=str,
|
||||
default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt",
|
||||
type=str,
|
||||
default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument(
|
||||
"--negative_prompt",
|
||||
type=str,
|
||||
default=(
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=[
|
||||
"left", "right", "up", "down", "zoom_in",
|
||||
"zoom_out", "clockwise", "counterclockwise", "none"
|
||||
])
|
||||
parser.add_argument("--movement_distance", type=float, default=0.3)
|
||||
parser.add_argument("--camera_rotation",
|
||||
type=str,
|
||||
default="center_facing",
|
||||
choices=[
|
||||
"center_facing", "no_rotation",
|
||||
"trajectory_aligned"
|
||||
])
|
||||
parser.add_argument("--height", type=int, default=704)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=121)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=35)
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0)
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
video = generator.generate_video(
|
||||
args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
image_path=args.image_path,
|
||||
trajectory_type=args.trajectory,
|
||||
movement_distance=args.movement_distance,
|
||||
camera_rotation=args.camera_rotation,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
fps=24,
|
||||
seed=args.seed,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -10,6 +10,35 @@ set -ex
|
||||
|
||||
echo "Building fastvideo-kernel..."
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Neutralise conda-injected compiler toolchains.
|
||||
#
|
||||
# Conda compiler packages (gcc_linux-aarch64, gxx_linux-64, etc.) set
|
||||
# CMAKE_ARGS, CFLAGS, CXXFLAGS, and LDFLAGS on activation. When multiple
|
||||
# toolchains are installed the variables can reference a *cross*-compiler
|
||||
# that doesn't match the host (e.g. aarch64-conda-linux-gnu-c++ on x86_64).
|
||||
# Even when the correct toolchain is active, the flags it injects
|
||||
# (-march=nocona, -mtune=haswell, …) can conflict with nvcc's host-compiler
|
||||
# expectations. Clear them so CMake discovers the system compiler instead.
|
||||
# ---------------------------------------------------------------------------
|
||||
if [[ -n "${CONDA_PREFIX:-}" ]]; then
|
||||
_need_clean=0
|
||||
# Detect conda cross-compiler that doesn't match the host.
|
||||
_host_arch="$(uname -m)"
|
||||
if [[ "${CXX:-}" == *"conda"* ]] || [[ "${CC:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if [[ "${CMAKE_ARGS:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if (( _need_clean )); then
|
||||
echo "NOTE: Clearing conda-injected compiler settings (CC/CXX/CMAKE_ARGS/CFLAGS/...)"
|
||||
echo " to use the system compiler for CUDA extension builds."
|
||||
unset CC CXX CMAKE_ARGS CFLAGS CXXFLAGS LDFLAGS
|
||||
fi
|
||||
unset _need_clean _host_arch
|
||||
fi
|
||||
|
||||
# Ensure submodules are initialized if needed (tk)
|
||||
git submodule update --init --recursive
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0",
|
||||
"triton>=2.0.0; sys_platform == 'linux'",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ComponentConfig,
|
||||
ContinuationState,
|
||||
EngineConfig,
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
PlannedStage,
|
||||
QuantizationConfig,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
ServerConfig,
|
||||
)
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.parser import (
|
||||
config_to_dict,
|
||||
load_config,
|
||||
load_raw_config,
|
||||
load_run_config,
|
||||
load_serve_config,
|
||||
parse_config,
|
||||
)
|
||||
from fastvideo.api.results import GenerationResult
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"ComponentConfig",
|
||||
"ContinuationState",
|
||||
"ConfigValidationError",
|
||||
"EngineConfig",
|
||||
"GenerationResult",
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"apply_overrides",
|
||||
"config_to_dict",
|
||||
"load_config",
|
||||
"load_raw_config",
|
||||
"load_run_config",
|
||||
"load_serve_config",
|
||||
"parse_cli_overrides",
|
||||
"parse_config",
|
||||
]
|
||||
@@ -0,0 +1,507 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from dataclasses import fields, is_dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.parser import config_to_dict, load_raw_config, parse_config
|
||||
from fastvideo.api.request_metadata import (
|
||||
EXPLICIT_REQUEST_ATTR,
|
||||
bind_generation_request_raw,
|
||||
refresh_generation_request_raw,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
RequestRuntimeConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
|
||||
_SAMPLING_FIELD_NAMES = {field.name for field in fields(SamplingConfig)}
|
||||
_RUNTIME_FIELD_NAMES = {field.name for field in fields(RequestRuntimeConfig)}
|
||||
_OUTPUT_FIELD_NAMES = {field.name for field in fields(OutputConfig)}
|
||||
_MISSING = object()
|
||||
_LEGACY_REQUEST_ALIASES = {
|
||||
"neg_prompt": "negative_prompt",
|
||||
}
|
||||
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
|
||||
"embedded_cfg_scale",
|
||||
})
|
||||
|
||||
|
||||
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
|
||||
if isinstance(config, GeneratorConfig):
|
||||
return config
|
||||
return parse_config(GeneratorConfig, config)
|
||||
|
||||
|
||||
def load_generator_config_from_file(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> GeneratorConfig:
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
|
||||
if _looks_like_run_or_serve_config(raw):
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(GeneratorConfig, raw["generator"])
|
||||
|
||||
if normalized_overrides:
|
||||
adjusted = normalized_overrides
|
||||
if all(key.startswith("generator.") for key in adjusted):
|
||||
adjusted = {key[len("generator."):]: value for key, value in adjusted.items()}
|
||||
raw = apply_overrides(raw, adjusted)
|
||||
|
||||
return parse_config(GeneratorConfig, raw)
|
||||
|
||||
|
||||
def legacy_from_pretrained_to_config(
|
||||
model_path: str,
|
||||
kwargs: Mapping[str, Any],
|
||||
) -> GeneratorConfig:
|
||||
raw: dict[str, Any] = {"model_path": model_path}
|
||||
engine: dict[str, Any] = {}
|
||||
parallelism: dict[str, Any] = {}
|
||||
offload: dict[str, Any] = {}
|
||||
compile_config: dict[str, Any] = {}
|
||||
pipeline: dict[str, Any] = {}
|
||||
components: dict[str, Any] = {}
|
||||
quantization: dict[str, Any] = {}
|
||||
experimental: dict[str, Any] = {}
|
||||
|
||||
for key, value in kwargs.items():
|
||||
if key == "revision":
|
||||
raw["revision"] = value
|
||||
elif key == "trust_remote_code":
|
||||
raw["trust_remote_code"] = value
|
||||
elif key == "num_gpus":
|
||||
engine["num_gpus"] = value
|
||||
elif key == "distributed_executor_backend":
|
||||
engine["execution_backend"] = value
|
||||
elif key in {"tp_size", "sp_size", "hsdp_replicate_dim", "hsdp_shard_dim", "dist_timeout"}:
|
||||
parallelism[key] = value
|
||||
elif key == "dit_cpu_offload":
|
||||
offload["dit"] = value
|
||||
elif key == "dit_layerwise_offload":
|
||||
offload["dit_layerwise"] = value
|
||||
elif key == "text_encoder_cpu_offload":
|
||||
offload["text_encoder"] = value
|
||||
elif key == "image_encoder_cpu_offload":
|
||||
offload["image_encoder"] = value
|
||||
elif key == "vae_cpu_offload":
|
||||
offload["vae"] = value
|
||||
elif key == "pin_cpu_memory":
|
||||
offload["pin_cpu_memory"] = value
|
||||
elif key == "enable_torch_compile":
|
||||
compile_config["enabled"] = value
|
||||
elif key == "torch_compile_kwargs":
|
||||
compile_config["kwargs"] = deepcopy(value)
|
||||
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
|
||||
engine[key] = value
|
||||
elif key == "override_text_encoder_quant":
|
||||
quantization["text_encoder_quant"] = value
|
||||
elif key == "workload_type":
|
||||
pipeline["workload_type"] = value
|
||||
elif key == "lora_path":
|
||||
components["lora_path"] = value
|
||||
elif key == "override_pipeline_cls_name":
|
||||
components["override_pipeline_cls_name"] = value
|
||||
elif key == "override_transformer_cls_name":
|
||||
components["override_transformer_cls_name"] = value
|
||||
elif key == "pipeline_config":
|
||||
if isinstance(value, str):
|
||||
components["pipeline_config_path"] = value
|
||||
else:
|
||||
experimental[key] = deepcopy(value)
|
||||
elif key == "override_text_encoder_safetensors":
|
||||
components["text_encoder_weights"] = value
|
||||
elif key == "init_weights_from_safetensors":
|
||||
components["transformer_weights"] = value
|
||||
elif key == "init_weights_from_safetensors_2":
|
||||
components["transformer_2_weights"] = value
|
||||
else:
|
||||
experimental[key] = deepcopy(value)
|
||||
|
||||
if parallelism:
|
||||
engine["parallelism"] = parallelism
|
||||
if offload:
|
||||
engine["offload"] = offload
|
||||
if compile_config:
|
||||
engine["compile"] = compile_config
|
||||
if quantization:
|
||||
engine["quantization"] = quantization
|
||||
if engine:
|
||||
raw["engine"] = engine
|
||||
|
||||
if components:
|
||||
pipeline["components"] = components
|
||||
if experimental:
|
||||
pipeline["experimental"] = experimental
|
||||
if pipeline:
|
||||
raw["pipeline"] = pipeline
|
||||
|
||||
return parse_config(GeneratorConfig, raw)
|
||||
|
||||
|
||||
def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, Any], ) -> FastVideoArgs:
|
||||
normalized = normalize_generator_config(config)
|
||||
unsupported = []
|
||||
if normalized.pipeline.profile is not None:
|
||||
unsupported.append("pipeline.profile")
|
||||
if normalized.pipeline.profile_version is not None:
|
||||
unsupported.append("pipeline.profile_version")
|
||||
if normalized.pipeline.components.config_root is not None:
|
||||
unsupported.append("pipeline.components.config_root")
|
||||
if normalized.pipeline.components.vae_weights is not None:
|
||||
unsupported.append("pipeline.components.vae_weights")
|
||||
if normalized.pipeline.components.upsampler_weights is not None:
|
||||
unsupported.append("pipeline.components.upsampler_weights")
|
||||
if unsupported:
|
||||
joined = ", ".join(unsupported)
|
||||
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
|
||||
|
||||
engine = normalized.engine
|
||||
kwargs: dict[str, Any] = {
|
||||
"model_path": normalized.model_path,
|
||||
"revision": normalized.revision,
|
||||
"trust_remote_code": normalized.trust_remote_code,
|
||||
"num_gpus": engine.num_gpus,
|
||||
"distributed_executor_backend": engine.execution_backend,
|
||||
"tp_size": engine.parallelism.tp_size,
|
||||
"sp_size": engine.parallelism.sp_size,
|
||||
"hsdp_replicate_dim": engine.parallelism.hsdp_replicate_dim,
|
||||
"hsdp_shard_dim": engine.parallelism.hsdp_shard_dim,
|
||||
"dist_timeout": engine.parallelism.dist_timeout,
|
||||
"dit_cpu_offload": engine.offload.dit,
|
||||
"dit_layerwise_offload": engine.offload.dit_layerwise,
|
||||
"text_encoder_cpu_offload": engine.offload.text_encoder,
|
||||
"image_encoder_cpu_offload": engine.offload.image_encoder,
|
||||
"vae_cpu_offload": engine.offload.vae,
|
||||
"pin_cpu_memory": engine.offload.pin_cpu_memory,
|
||||
"enable_torch_compile": engine.compile.enabled,
|
||||
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
|
||||
"enable_stage_verification": engine.enable_stage_verification,
|
||||
"use_fsdp_inference": engine.use_fsdp_inference,
|
||||
"disable_autocast": engine.disable_autocast,
|
||||
}
|
||||
if normalized.pipeline.workload_type is not None:
|
||||
kwargs["workload_type"] = normalized.pipeline.workload_type
|
||||
|
||||
quantization = engine.quantization
|
||||
if quantization is not None and quantization.text_encoder_quant is not None:
|
||||
kwargs["override_text_encoder_quant"] = quantization.text_encoder_quant
|
||||
if quantization is not None and quantization.transformer_quant is not None:
|
||||
kwargs["transformer_quant"] = quantization.transformer_quant
|
||||
|
||||
components = normalized.pipeline.components
|
||||
if components.pipeline_config_path is not None:
|
||||
kwargs["pipeline_config"] = components.pipeline_config_path
|
||||
if components.lora_path is not None:
|
||||
kwargs["lora_path"] = components.lora_path
|
||||
if components.override_pipeline_cls_name is not None:
|
||||
kwargs["override_pipeline_cls_name"] = components.override_pipeline_cls_name
|
||||
if components.override_transformer_cls_name is not None:
|
||||
kwargs["override_transformer_cls_name"] = components.override_transformer_cls_name
|
||||
if components.text_encoder_weights is not None:
|
||||
kwargs["override_text_encoder_safetensors"] = components.text_encoder_weights
|
||||
if components.transformer_weights is not None:
|
||||
kwargs["init_weights_from_safetensors"] = components.transformer_weights
|
||||
if components.transformer_2_weights is not None:
|
||||
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
|
||||
|
||||
kwargs.update(deepcopy(normalized.pipeline.profile_overrides))
|
||||
kwargs.update(deepcopy(normalized.pipeline.experimental))
|
||||
return FastVideoArgs.from_kwargs(**kwargs)
|
||||
|
||||
|
||||
def normalize_generation_request(request: GenerationRequest | Mapping[str, Any], ) -> GenerationRequest:
|
||||
normalized = (request if isinstance(request, GenerationRequest) else parse_config(GenerationRequest, request))
|
||||
|
||||
if hasattr(normalized, EXPLICIT_REQUEST_ATTR):
|
||||
refresh_generation_request_raw(normalized)
|
||||
else:
|
||||
bind_generation_request_raw(normalized, _serialize_generation_request(normalized))
|
||||
return normalized
|
||||
|
||||
|
||||
def legacy_generate_call_to_request(
|
||||
prompt: str | None,
|
||||
sampling_param: SamplingParam | None,
|
||||
*,
|
||||
mouse_cond: Any | None = None,
|
||||
keyboard_cond: Any | None = None,
|
||||
grid_sizes: Any | None = None,
|
||||
legacy_kwargs: Mapping[str, Any] | None = None,
|
||||
) -> GenerationRequest:
|
||||
raw = _sampling_param_to_request_raw(sampling_param)
|
||||
if prompt is not None:
|
||||
raw["prompt"] = prompt
|
||||
|
||||
for key, value in (legacy_kwargs or {}).items():
|
||||
_apply_request_field(raw, key, value)
|
||||
|
||||
if mouse_cond is not None:
|
||||
raw.setdefault("inputs", {})["mouse_cond"] = mouse_cond
|
||||
if keyboard_cond is not None:
|
||||
raw.setdefault("inputs", {})["keyboard_cond"] = keyboard_cond
|
||||
if grid_sizes is not None:
|
||||
raw.setdefault("inputs", {})["grid_sizes"] = grid_sizes
|
||||
|
||||
normalized = parse_config(GenerationRequest, raw)
|
||||
bind_generation_request_raw(normalized, raw)
|
||||
return normalized
|
||||
|
||||
|
||||
def request_to_sampling_param(
|
||||
request: GenerationRequest,
|
||||
*,
|
||||
model_path: str,
|
||||
) -> SamplingParam:
|
||||
if request.plan is not None:
|
||||
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
|
||||
if request.state is not None:
|
||||
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
updates = _explicit_request_updates(request)
|
||||
|
||||
for key, value in updates.items():
|
||||
if hasattr(sampling_param, key):
|
||||
setattr(sampling_param, key, deepcopy(value))
|
||||
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS or _is_supported_as_default_only(key, value):
|
||||
continue
|
||||
else:
|
||||
raise ValueError(f"Request field {key!r} is not supported by sampling params for {model_path}")
|
||||
|
||||
sampling_param.__post_init__()
|
||||
sampling_param.check_sampling_param()
|
||||
return sampling_param
|
||||
|
||||
|
||||
def expand_request_prompt_batch(request: GenerationRequest, ) -> list[GenerationRequest]:
|
||||
if not isinstance(request.prompt, list):
|
||||
return [request]
|
||||
|
||||
requests: list[GenerationRequest] = []
|
||||
for index, prompt in enumerate(request.prompt):
|
||||
single_request = deepcopy(request)
|
||||
single_request.prompt = prompt
|
||||
_fan_out_batched_input_value(request, single_request, "image_path", index)
|
||||
_fan_out_batched_input_value(request, single_request, "video_path", index)
|
||||
_fan_out_explicit_request_metadata(request, single_request, index, prompt)
|
||||
requests.append(single_request)
|
||||
return requests
|
||||
|
||||
|
||||
def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
return isinstance(raw.get("generator"), Mapping)
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
def _sampling_param_to_request_raw(sampling_param: SamplingParam | None, ) -> dict[str, Any]:
|
||||
if sampling_param is None:
|
||||
return {}
|
||||
|
||||
raw: dict[str, Any] = {}
|
||||
for key, value in shallow_asdict(sampling_param).items():
|
||||
if key == "prompt":
|
||||
continue
|
||||
_apply_request_field(raw, key, deepcopy(value))
|
||||
return raw
|
||||
|
||||
|
||||
def _apply_request_field(
|
||||
raw: dict[str, Any],
|
||||
key: str,
|
||||
value: Any,
|
||||
) -> None:
|
||||
key = _LEGACY_REQUEST_ALIASES.get(key, key)
|
||||
if key == "negative_prompt":
|
||||
raw["negative_prompt"] = value
|
||||
return
|
||||
if key in _INPUT_FIELD_NAMES:
|
||||
raw.setdefault("inputs", {})[key] = value
|
||||
return
|
||||
if key in _SAMPLING_FIELD_NAMES:
|
||||
raw.setdefault("sampling", {})[key] = value
|
||||
return
|
||||
if key in _RUNTIME_FIELD_NAMES:
|
||||
raw.setdefault("runtime", {})[key] = value
|
||||
return
|
||||
if key in _OUTPUT_FIELD_NAMES:
|
||||
raw.setdefault("output", {})[key] = value
|
||||
return
|
||||
raw.setdefault("extensions", {})[key] = value
|
||||
|
||||
|
||||
def request_to_pipeline_overrides(request: GenerationRequest) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {}
|
||||
for key, value in _explicit_request_updates(request).items():
|
||||
if key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
|
||||
overrides[key] = deepcopy(value)
|
||||
return overrides
|
||||
|
||||
|
||||
def _explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
|
||||
raw = getattr(request, EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
raw = _serialize_generation_request(request)
|
||||
|
||||
return _extract_request_updates(raw)
|
||||
|
||||
|
||||
def _extract_request_updates(raw: Mapping[str, Any]) -> dict[str, Any]:
|
||||
updates: dict[str, Any] = {}
|
||||
if "negative_prompt" in raw:
|
||||
updates["negative_prompt"] = deepcopy(raw["negative_prompt"])
|
||||
|
||||
for section_name in ("inputs", "sampling", "runtime", "output"):
|
||||
section = raw.get(section_name)
|
||||
if not isinstance(section, Mapping):
|
||||
continue
|
||||
for key, value in section.items():
|
||||
updates[key] = deepcopy(value)
|
||||
|
||||
stage_overrides = raw.get("stage_overrides")
|
||||
if stage_overrides:
|
||||
updates.update(_flatten_stage_overrides(stage_overrides))
|
||||
|
||||
extensions = raw.get("extensions")
|
||||
if isinstance(extensions, Mapping):
|
||||
for key, value in extensions.items():
|
||||
updates[key] = deepcopy(value)
|
||||
|
||||
return updates
|
||||
|
||||
|
||||
def _flatten_stage_overrides(stage_overrides: Any) -> dict[str, Any]:
|
||||
if not isinstance(stage_overrides, Mapping):
|
||||
raise ValueError("GenerationRequest.stage_overrides must be a mapping")
|
||||
|
||||
flattened: dict[str, Any] = {}
|
||||
for stage_name, overrides in stage_overrides.items():
|
||||
if not isinstance(overrides, Mapping):
|
||||
raise ValueError(f"GenerationRequest.stage_overrides.{stage_name} must be a mapping")
|
||||
for key, value in overrides.items():
|
||||
if key in flattened and flattened[key] != value:
|
||||
raise ValueError(f"Conflicting stage override for {key!r} across stages")
|
||||
flattened[key] = deepcopy(value)
|
||||
return flattened
|
||||
|
||||
|
||||
def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
|
||||
return deepcopy(config_to_dict(request))
|
||||
|
||||
|
||||
def _fan_out_batched_input_value(
|
||||
source_request: GenerationRequest,
|
||||
target_request: GenerationRequest,
|
||||
field_name: str,
|
||||
index: int,
|
||||
) -> None:
|
||||
value = getattr(source_request.inputs, field_name)
|
||||
if not isinstance(value, list):
|
||||
return
|
||||
_validate_batched_input_length(source_request.prompt, value, field_name)
|
||||
setattr(target_request.inputs, field_name, deepcopy(value[index]))
|
||||
|
||||
|
||||
def _fan_out_explicit_request_metadata(
|
||||
source_request: GenerationRequest,
|
||||
target_request: GenerationRequest,
|
||||
index: int,
|
||||
prompt: str,
|
||||
) -> None:
|
||||
raw = getattr(source_request, EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
return
|
||||
|
||||
raw = deepcopy(raw)
|
||||
raw["prompt"] = prompt
|
||||
inputs = raw.get("inputs")
|
||||
if isinstance(inputs, dict):
|
||||
for field_name in ("image_path", "video_path"):
|
||||
value = inputs.get(field_name)
|
||||
if isinstance(value, list):
|
||||
_validate_batched_input_length(source_request.prompt, value, field_name)
|
||||
inputs[field_name] = deepcopy(value[index])
|
||||
|
||||
setattr(target_request, EXPLICIT_REQUEST_ATTR, raw)
|
||||
|
||||
|
||||
def _validate_batched_input_length(
|
||||
prompts: str | list[str] | None,
|
||||
values: list[Any],
|
||||
field_name: str,
|
||||
) -> None:
|
||||
if not isinstance(prompts, list):
|
||||
return
|
||||
if len(values) != len(prompts):
|
||||
raise ValueError(f"GenerationRequest.inputs.{field_name} must have the same length as request.prompt")
|
||||
|
||||
|
||||
def _is_supported_as_default_only(key: str, value: Any) -> bool:
|
||||
default_value = _DEFAULT_REQUEST_UPDATES.get(key, _MISSING)
|
||||
return default_value is not _MISSING and _values_equal(value, default_value)
|
||||
|
||||
|
||||
def _collect_non_default_fields(
|
||||
value: Any,
|
||||
default: Any,
|
||||
) -> dict[str, Any]:
|
||||
if not (is_dataclass(value) and is_dataclass(default)):
|
||||
return {}
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
for field in fields(value):
|
||||
current = getattr(value, field.name)
|
||||
default_value = getattr(default, field.name)
|
||||
if is_dataclass(current) and is_dataclass(default_value):
|
||||
nested = _collect_non_default_fields(current, default_value)
|
||||
if nested:
|
||||
result[field.name] = nested
|
||||
continue
|
||||
if not _values_equal(current, default_value):
|
||||
result[field.name] = deepcopy(current)
|
||||
return result
|
||||
|
||||
|
||||
def _values_equal(left: Any, right: Any) -> bool:
|
||||
if left is right:
|
||||
return True
|
||||
try:
|
||||
return bool(left == right)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
_DEFAULT_REQUEST_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
__all__ = [
|
||||
"generator_config_to_fastvideo_args",
|
||||
"legacy_from_pretrained_to_config",
|
||||
"legacy_generate_call_to_request",
|
||||
"load_generator_config_from_file",
|
||||
"normalize_generation_request",
|
||||
"normalize_generator_config",
|
||||
"request_to_pipeline_overrides",
|
||||
"request_to_sampling_param",
|
||||
]
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class ConfigValidationError(ValueError):
|
||||
"""Validation error that keeps track of the nested config path."""
|
||||
|
||||
def __init__(self, path: str, message: str):
|
||||
self.path = path
|
||||
self.message = message
|
||||
super().__init__(str(self))
|
||||
|
||||
def __str__(self) -> str:
|
||||
if self.path:
|
||||
return f"{self.path}: {self.message}"
|
||||
return self.message
|
||||
@@ -0,0 +1,101 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
from collections.abc import Mapping
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
|
||||
|
||||
def parse_cli_overrides(overrides: list[str]) -> dict[str, Any]:
|
||||
"""Parse ``--dotted.key value`` style overrides into a flat mapping."""
|
||||
parsed: dict[str, Any] = {}
|
||||
index = 0
|
||||
while index < len(overrides):
|
||||
token = overrides[index]
|
||||
if not token.startswith("--"):
|
||||
raise ValueError(f"Expected --dotted.key, got {token!r}")
|
||||
|
||||
key = token[2:]
|
||||
if not key:
|
||||
raise ValueError("Override key cannot be empty")
|
||||
|
||||
if "=" in key:
|
||||
key, raw_value = key.split("=", 1)
|
||||
else:
|
||||
index += 1
|
||||
if index >= len(overrides):
|
||||
raise ValueError(f"Missing value for override {token!r}")
|
||||
raw_value = overrides[index]
|
||||
|
||||
parsed[_normalize_override_key(key)] = _cast_override_value(raw_value)
|
||||
index += 1
|
||||
|
||||
return parsed
|
||||
|
||||
|
||||
def apply_overrides(config: Mapping[str, Any], overrides: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of ``config`` with dotted-key overrides applied."""
|
||||
merged = deepcopy(dict(config))
|
||||
for dotted_key, value in overrides.items():
|
||||
_apply_single_override(merged, dotted_key, value)
|
||||
return merged
|
||||
|
||||
|
||||
def _apply_single_override(config: dict[str, Any], dotted_key: str, value: Any) -> None:
|
||||
parts = dotted_key.split(".")
|
||||
if not all(parts):
|
||||
raise ValueError(f"Invalid override path {dotted_key!r}")
|
||||
|
||||
cursor = config
|
||||
for depth, part in enumerate(parts[:-1]):
|
||||
existing = cursor.get(part)
|
||||
if existing is None:
|
||||
existing = {}
|
||||
cursor[part] = existing
|
||||
elif not isinstance(existing, dict):
|
||||
raise ConfigValidationError(
|
||||
".".join(parts[:depth + 1]),
|
||||
"cannot apply nested override through a non-mapping value",
|
||||
)
|
||||
cursor = existing
|
||||
|
||||
cursor[parts[-1]] = value
|
||||
|
||||
|
||||
def _cast_override_value(raw: str) -> Any:
|
||||
lowered = raw.lower()
|
||||
if lowered == "true":
|
||||
return True
|
||||
if lowered == "false":
|
||||
return False
|
||||
if lowered in {"none", "null"}:
|
||||
return None
|
||||
|
||||
try:
|
||||
return int(raw)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
try:
|
||||
return float(raw)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if raw.startswith("[") or raw.startswith("{"):
|
||||
try:
|
||||
return yaml.safe_load(raw)
|
||||
except yaml.YAMLError:
|
||||
pass
|
||||
|
||||
return raw
|
||||
|
||||
|
||||
def _normalize_override_key(key: str) -> str:
|
||||
return key.replace("-", "_")
|
||||
|
||||
|
||||
__all__ = ["apply_overrides", "parse_cli_overrides"]
|
||||
@@ -0,0 +1,332 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import json
|
||||
import types
|
||||
from pathlib import Path
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Literal, TypeVar, Union, get_args, get_origin, get_type_hints
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.request_metadata import (
|
||||
bind_generation_request_raw,
|
||||
bind_run_config_raw,
|
||||
bind_serve_config_raw,
|
||||
)
|
||||
from fastvideo.api.schema import GenerationRequest, RunConfig, ServeConfig
|
||||
|
||||
T = TypeVar("T")
|
||||
_UNION_ORIGINS = {types.UnionType, Union}
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class _DataclassSpec:
|
||||
cls: type[Any]
|
||||
type_hints: dict[str, Any]
|
||||
fields_by_name: dict[str, dataclasses.Field[Any]]
|
||||
|
||||
|
||||
def parse_config(config_type: type[T], raw: Mapping[str, Any] | T) -> T:
|
||||
"""Parse a nested mapping into a typed inference config object."""
|
||||
if isinstance(raw, config_type):
|
||||
return raw
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError("", f"expected mapping for {config_type.__name__}")
|
||||
parsed = _SchemaParser().parse_dataclass(config_type, raw, "")
|
||||
if config_type is GenerationRequest:
|
||||
return bind_generation_request_raw(parsed, raw)
|
||||
if config_type is RunConfig:
|
||||
return bind_run_config_raw(parsed, raw)
|
||||
if config_type is ServeConfig:
|
||||
return bind_serve_config_raw(parsed, raw)
|
||||
return parsed
|
||||
|
||||
|
||||
def config_to_dict(config: Any) -> Any:
|
||||
"""Serialize a typed config object into plain Python containers."""
|
||||
if dataclasses.is_dataclass(config) and not isinstance(config, type):
|
||||
return {field.name: config_to_dict(getattr(config, field.name)) for field in dataclasses.fields(config)}
|
||||
if isinstance(config, list):
|
||||
return [config_to_dict(item) for item in config]
|
||||
if isinstance(config, dict):
|
||||
return {key: config_to_dict(value) for key, value in config.items()}
|
||||
return config
|
||||
|
||||
|
||||
def load_config(
|
||||
config_type: type[T],
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> T:
|
||||
"""Load a typed config object from YAML or JSON."""
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(config_type, raw)
|
||||
|
||||
|
||||
def load_run_config(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> RunConfig:
|
||||
return load_config(RunConfig, path, overrides)
|
||||
|
||||
|
||||
def load_serve_config(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> ServeConfig:
|
||||
return load_config(ServeConfig, path, overrides)
|
||||
|
||||
|
||||
def load_raw_config(path: str | Path) -> dict[str, Any]:
|
||||
config_path = Path(path)
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||
|
||||
with config_path.open(encoding="utf-8") as handle:
|
||||
raw = _load_raw_mapping(handle, config_path)
|
||||
|
||||
if raw is None:
|
||||
return {}
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError("", f"{config_path} must contain a top-level mapping")
|
||||
return dict(raw)
|
||||
|
||||
|
||||
def _load_raw_mapping(handle: Any, config_path: Path) -> Any:
|
||||
suffix = config_path.suffix.lower()
|
||||
if suffix in {".yaml", ".yml"}:
|
||||
return yaml.safe_load(handle)
|
||||
if suffix == ".json":
|
||||
return json.load(handle)
|
||||
raise ValueError(f"Unsupported config file format: {config_path}")
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
class _SchemaParser:
|
||||
|
||||
def parse_dataclass(
|
||||
self,
|
||||
config_type: type[T],
|
||||
raw: Mapping[str, Any],
|
||||
path: str,
|
||||
) -> T:
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError(path, f"expected mapping for {config_type.__name__}")
|
||||
|
||||
spec = _get_dataclass_spec(config_type)
|
||||
self._validate_keys(raw, spec, path)
|
||||
|
||||
values: dict[str, Any] = {}
|
||||
for name, field in spec.fields_by_name.items():
|
||||
field_path = _join_path(path, name)
|
||||
if name in raw:
|
||||
values[name] = self.parse_value(spec.type_hints[name], raw[name], field_path)
|
||||
continue
|
||||
if _field_is_required(field):
|
||||
raise ConfigValidationError(field_path, "missing required field")
|
||||
|
||||
return config_type(**values)
|
||||
|
||||
def parse_value(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if annotation is Any:
|
||||
return value
|
||||
|
||||
origin = get_origin(annotation)
|
||||
if origin in _UNION_ORIGINS:
|
||||
return self._parse_union(annotation, value, path)
|
||||
if origin is Literal:
|
||||
return self._parse_literal(annotation, value, path)
|
||||
if origin is list:
|
||||
return self._parse_list(annotation, value, path)
|
||||
if origin is dict:
|
||||
return self._parse_dict(annotation, value, path)
|
||||
if origin is tuple:
|
||||
return self._parse_tuple(annotation, value, path)
|
||||
if isinstance(annotation, type) and dataclasses.is_dataclass(annotation):
|
||||
return self.parse_dataclass(annotation, value, path)
|
||||
|
||||
scalar_parser = _SCALAR_PARSERS.get(annotation)
|
||||
if scalar_parser is not None:
|
||||
return scalar_parser(value, path)
|
||||
|
||||
return self._parse_instance(annotation, value, path)
|
||||
|
||||
def _validate_keys(
|
||||
self,
|
||||
raw: Mapping[str, Any],
|
||||
spec: _DataclassSpec,
|
||||
path: str,
|
||||
) -> None:
|
||||
for key in raw:
|
||||
if not isinstance(key, str):
|
||||
raise ConfigValidationError(path, "expected mapping keys to be strings")
|
||||
if key not in spec.fields_by_name:
|
||||
raise ConfigValidationError(_join_path(path, key), "unknown field")
|
||||
|
||||
def _parse_union(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
candidates = [candidate for candidate in get_args(annotation) if candidate is not type(None)]
|
||||
if value is None and len(candidates) != len(get_args(annotation)):
|
||||
return None
|
||||
if len(candidates) == 1:
|
||||
return self.parse_value(candidates[0], value, path)
|
||||
|
||||
errors: list[str] = []
|
||||
for candidate in candidates:
|
||||
try:
|
||||
return self.parse_value(candidate, value, path)
|
||||
except ConfigValidationError as exc:
|
||||
errors.append(exc.message)
|
||||
|
||||
expected = ", ".join(_type_name(candidate) for candidate in candidates)
|
||||
detail = errors[0] if errors else f"expected one of ({expected})"
|
||||
raise ConfigValidationError(path, detail)
|
||||
|
||||
def _parse_literal(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
allowed = get_args(annotation)
|
||||
if value not in allowed:
|
||||
raise ConfigValidationError(path, f"expected one of {sorted(allowed)!r}")
|
||||
return value
|
||||
|
||||
def _parse_list(self, annotation: Any, value: Any, path: str) -> list[Any]:
|
||||
if not isinstance(value, list):
|
||||
raise ConfigValidationError(path, "expected list")
|
||||
item_type = get_args(annotation)[0] if get_args(annotation) else Any
|
||||
return [self.parse_value(item_type, item, f"{path}[{index}]") for index, item in enumerate(value)]
|
||||
|
||||
def _parse_dict(self, annotation: Any, value: Any, path: str) -> dict[Any, Any]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ConfigValidationError(path, "expected mapping")
|
||||
|
||||
key_type, value_type = (get_args(annotation) + (Any, Any))[:2]
|
||||
parsed: dict[Any, Any] = {}
|
||||
for key, item in value.items():
|
||||
parsed_key = self._parse_dict_key(key_type, key, path)
|
||||
item_path = _join_path(path, str(key))
|
||||
parsed[parsed_key] = self.parse_value(value_type, item, item_path)
|
||||
return parsed
|
||||
|
||||
def _parse_tuple(self, annotation: Any, value: Any, path: str) -> tuple[Any, ...]:
|
||||
if not isinstance(value, list | tuple):
|
||||
raise ConfigValidationError(path, "expected tuple")
|
||||
|
||||
item_types = get_args(annotation)
|
||||
if len(item_types) == 2 and item_types[1] is Ellipsis:
|
||||
return tuple(self.parse_value(item_types[0], item, f"{path}[{index}]") for index, item in enumerate(value))
|
||||
|
||||
if len(value) != len(item_types):
|
||||
raise ConfigValidationError(path, f"expected tuple of length {len(item_types)}")
|
||||
|
||||
return tuple(
|
||||
self.parse_value(item_type, item, f"{path}[{index}]")
|
||||
for index, (item_type, item) in enumerate(zip(item_types, value, strict=True)))
|
||||
|
||||
def _parse_dict_key(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if annotation is Any:
|
||||
return value
|
||||
if annotation is str:
|
||||
if not isinstance(value, str):
|
||||
raise ConfigValidationError(path, "expected string dictionary keys")
|
||||
return value
|
||||
if annotation is int:
|
||||
if not isinstance(value, int) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected integer dictionary keys")
|
||||
return value
|
||||
return value
|
||||
|
||||
def _parse_instance(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if isinstance(annotation, type) and not isinstance(value, annotation):
|
||||
raise ConfigValidationError(path, f"expected {annotation.__name__}")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_bool(value: Any, path: str) -> bool:
|
||||
if type(value) is not bool:
|
||||
raise ConfigValidationError(path, "expected bool")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_int(value: Any, path: str) -> int:
|
||||
if not isinstance(value, int) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected int")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_float(value: Any, path: str) -> float:
|
||||
if not isinstance(value, int | float) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected float")
|
||||
return float(value)
|
||||
|
||||
|
||||
def _parse_str(value: Any, path: str) -> str:
|
||||
if not isinstance(value, str):
|
||||
raise ConfigValidationError(path, "expected str")
|
||||
return value
|
||||
|
||||
|
||||
_SCALAR_PARSERS: dict[Any, Any] = {
|
||||
bool: _parse_bool,
|
||||
int: _parse_int,
|
||||
float: _parse_float,
|
||||
str: _parse_str,
|
||||
}
|
||||
|
||||
|
||||
def _field_is_required(field: dataclasses.Field[Any]) -> bool:
|
||||
return (field.default is dataclasses.MISSING and field.default_factory is dataclasses.MISSING)
|
||||
|
||||
|
||||
def _get_dataclass_spec(config_type: type[Any]) -> _DataclassSpec:
|
||||
spec = _DATACLASS_SPEC_CACHE.get(config_type)
|
||||
if spec is not None:
|
||||
return spec
|
||||
|
||||
spec = _DataclassSpec(
|
||||
cls=config_type,
|
||||
type_hints=get_type_hints(config_type),
|
||||
fields_by_name={field.name: field
|
||||
for field in dataclasses.fields(config_type)},
|
||||
)
|
||||
_DATACLASS_SPEC_CACHE[config_type] = spec
|
||||
return spec
|
||||
|
||||
|
||||
_DATACLASS_SPEC_CACHE: dict[type[Any], _DataclassSpec] = {}
|
||||
|
||||
|
||||
def _join_path(prefix: str, suffix: str) -> str:
|
||||
if not prefix:
|
||||
return suffix
|
||||
return f"{prefix}.{suffix}"
|
||||
|
||||
|
||||
def _type_name(annotation: Any) -> str:
|
||||
origin = get_origin(annotation)
|
||||
if origin is not None:
|
||||
return str(annotation)
|
||||
if hasattr(annotation, "__name__"):
|
||||
return annotation.__name__
|
||||
return str(annotation)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"config_to_dict",
|
||||
"load_config",
|
||||
"load_raw_config",
|
||||
"load_run_config",
|
||||
"load_serve_config",
|
||||
"parse_config",
|
||||
]
|
||||
@@ -0,0 +1,247 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track which GenerationRequest fields the user explicitly provided.
|
||||
|
||||
This module solves a specific problem: when translating a GenerationRequest into
|
||||
a legacy SamplingParam, we need to distinguish user-provided values (which
|
||||
should override model defaults) from schema defaults (which should NOT override
|
||||
model defaults).
|
||||
|
||||
The approach:
|
||||
1. At bind time, store the original raw dict and a baseline snapshot.
|
||||
2. Patch __setattr__ on tracked dataclass types to record dirty field paths.
|
||||
3. At access time, do a lazy 3-way merge: raw + baseline + current state,
|
||||
with dirty paths forcing inclusion even when current == baseline.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
import dataclasses
|
||||
from typing import Any, cast
|
||||
from collections.abc import Callable
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PlannedStage,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
EXPLICIT_REQUEST_ATTR = "_fastvideo_explicit_request"
|
||||
ORIGINAL_REQUEST_STATE_ATTR = "_fastvideo_original_request_state"
|
||||
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
|
||||
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
|
||||
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
|
||||
_DIRTY_PATHS_ATTR = "_fastvideo_dirty_paths"
|
||||
_TRACKED_REQUEST_TYPES = (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
SamplingConfig,
|
||||
RequestRuntimeConfig,
|
||||
OutputConfig,
|
||||
ContinuationState,
|
||||
PlannedStage,
|
||||
GenerationPlan,
|
||||
)
|
||||
|
||||
|
||||
def bind_generation_request_raw(
|
||||
request: GenerationRequest,
|
||||
raw: Mapping[str, Any] | None,
|
||||
) -> GenerationRequest:
|
||||
_ensure_request_tracking()
|
||||
# Disable dirty tracking during bind so tree walk doesn't record paths.
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, None)
|
||||
object.__setattr__(request, EXPLICIT_REQUEST_ATTR, deepcopy(dict(raw or {})))
|
||||
object.__setattr__(request, ORIGINAL_REQUEST_STATE_ATTR, _serialize_config(request))
|
||||
_set_tracking_roots(request, request, "")
|
||||
# Enable dirty tracking.
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, set())
|
||||
return request
|
||||
|
||||
|
||||
def bind_run_config_raw(
|
||||
config: RunConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> RunConfig:
|
||||
request_raw = raw.get("request")
|
||||
if isinstance(request_raw, Mapping):
|
||||
bind_generation_request_raw(config.request, request_raw)
|
||||
return config
|
||||
|
||||
|
||||
def bind_serve_config_raw(
|
||||
config: ServeConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> ServeConfig:
|
||||
default_request_raw = raw.get("default_request")
|
||||
if isinstance(default_request_raw, Mapping):
|
||||
bind_generation_request_raw(config.default_request, default_request_raw)
|
||||
elif "default_request" not in raw:
|
||||
bind_generation_request_raw(config.default_request, {})
|
||||
return config
|
||||
|
||||
|
||||
def refresh_generation_request_raw(request: GenerationRequest, ) -> dict[str, Any] | None:
|
||||
raw = getattr(request, EXPLICIT_REQUEST_ATTR, None)
|
||||
baseline = getattr(request, ORIGINAL_REQUEST_STATE_ATTR, None)
|
||||
if not isinstance(raw, Mapping) or not isinstance(baseline, Mapping):
|
||||
return None
|
||||
|
||||
dirty = getattr(request, _DIRTY_PATHS_ATTR, None) or frozenset()
|
||||
current = _serialize_config(request)
|
||||
merged = deepcopy(dict(raw))
|
||||
_merge_request_mutations(merged, dict(baseline), current, dirty)
|
||||
|
||||
object.__setattr__(request, EXPLICIT_REQUEST_ATTR, merged)
|
||||
object.__setattr__(request, ORIGINAL_REQUEST_STATE_ATTR, current)
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, set())
|
||||
return merged
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3-way merge: raw + baseline + current, with dirty-path forcing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_MISSING = object()
|
||||
|
||||
|
||||
def _merge_request_mutations(
|
||||
merged: dict[str, Any],
|
||||
baseline: Mapping[str, Any],
|
||||
current: Mapping[str, Any],
|
||||
dirty: frozenset[str] | set[str],
|
||||
path_prefix: str = "",
|
||||
force_dirty: bool = False,
|
||||
) -> None:
|
||||
# Remove keys that were deleted from the current state.
|
||||
for key in set(merged) | set(baseline):
|
||||
if key not in current:
|
||||
merged.pop(key, None)
|
||||
|
||||
for key in current:
|
||||
current_path = f"{path_prefix}.{key}" if path_prefix else key
|
||||
current_value = current[key]
|
||||
baseline_value = baseline.get(key, _MISSING)
|
||||
merged_value = merged.get(key, _MISSING)
|
||||
|
||||
# If this exact path was dirtied (e.g. whole section replaced),
|
||||
# propagate to all children.
|
||||
child_force = force_dirty or current_path in dirty
|
||||
|
||||
# Recurse into nested mappings.
|
||||
if isinstance(current_value, Mapping) and isinstance(baseline_value, Mapping):
|
||||
nested = (deepcopy(dict(merged_value)) if isinstance(merged_value, Mapping) else {})
|
||||
_merge_request_mutations(
|
||||
nested,
|
||||
baseline_value,
|
||||
current_value,
|
||||
dirty,
|
||||
current_path,
|
||||
child_force,
|
||||
)
|
||||
if nested:
|
||||
merged[key] = nested
|
||||
else:
|
||||
merged.pop(key, None)
|
||||
continue
|
||||
|
||||
# A field is explicitly set if:
|
||||
# - it's new (not in baseline),
|
||||
# - it changed from baseline,
|
||||
# - its path was touched by __setattr__ (dirty), or
|
||||
# - an ancestor path was dirty (whole section replaced).
|
||||
is_dirty = child_force or current_path in dirty
|
||||
if baseline_value is _MISSING or current_value != baseline_value or is_dirty:
|
||||
merged[key] = deepcopy(current_value)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# __setattr__ patching for dirty-path recording
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ensure_request_tracking() -> None:
|
||||
for config_type in _TRACKED_REQUEST_TYPES:
|
||||
_patch_tracking_setattr(config_type)
|
||||
|
||||
|
||||
def _patch_tracking_setattr(config_type: type[Any]) -> None:
|
||||
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
|
||||
return
|
||||
|
||||
original_setattr = cast(
|
||||
Callable[[Any, str, Any], None],
|
||||
config_type.__setattr__,
|
||||
)
|
||||
field_names = {field.name for field in dataclasses.fields(config_type)}
|
||||
|
||||
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
|
||||
if name.startswith("_fastvideo_") or name not in field_names:
|
||||
original_setattr(self, name, value)
|
||||
return
|
||||
|
||||
root = getattr(self, _TRACKING_ROOT_ATTR, None)
|
||||
if root is not None:
|
||||
dirty = getattr(root, _DIRTY_PATHS_ATTR, None)
|
||||
if isinstance(dirty, set):
|
||||
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
|
||||
path = f"{prefix}.{name}" if prefix else name
|
||||
dirty.add(path)
|
||||
|
||||
original_setattr(self, name, value)
|
||||
|
||||
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
|
||||
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tree walk to set tracking root/path on nested dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _set_tracking_roots(
|
||||
root: GenerationRequest,
|
||||
obj: Any,
|
||||
prefix: str,
|
||||
) -> None:
|
||||
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
|
||||
return
|
||||
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
|
||||
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
|
||||
for field in dataclasses.fields(obj):
|
||||
child = getattr(obj, field.name)
|
||||
child_path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
if dataclasses.is_dataclass(child) and not isinstance(child, type):
|
||||
_set_tracking_roots(root, child, child_path)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialization helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _serialize_config(config: Any) -> Any:
|
||||
if dataclasses.is_dataclass(config) and not isinstance(config, type):
|
||||
return {field.name: _serialize_config(getattr(config, field.name)) for field in dataclasses.fields(config)}
|
||||
if isinstance(config, list):
|
||||
return [_serialize_config(item) for item in config]
|
||||
if isinstance(config, dict):
|
||||
return {key: _serialize_config(value) for key, value in config.items()}
|
||||
return deepcopy(config)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EXPLICIT_REQUEST_ATTR",
|
||||
"ORIGINAL_REQUEST_STATE_ATTR",
|
||||
"bind_generation_request_raw",
|
||||
"bind_run_config_raw",
|
||||
"bind_serve_config_raw",
|
||||
"refresh_generation_request_raw",
|
||||
]
|
||||
@@ -0,0 +1,101 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from collections.abc import Mapping
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationResult:
|
||||
prompt: str | None = None
|
||||
prompt_index: int | None = None
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
trajectory: Any | None = None
|
||||
trajectory_timesteps: Any | None = None
|
||||
trajectory_decoded: Any | None = None
|
||||
video_path: str | None = None
|
||||
peak_memory_mb: float | None = None
|
||||
state: ContinuationState | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_legacy_result(
|
||||
cls,
|
||||
result: Mapping[str, Any],
|
||||
) -> GenerationResult:
|
||||
prompt = result.get("prompt")
|
||||
if prompt is None:
|
||||
prompt = result.get("prompts")
|
||||
|
||||
extra = {
|
||||
key: value
|
||||
for key, value in result.items() if key not in {
|
||||
"prompt",
|
||||
"prompt_index",
|
||||
"prompts",
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
"trajectory",
|
||||
"trajectory_timesteps",
|
||||
"trajectory_decoded",
|
||||
"video_path",
|
||||
"peak_memory_mb",
|
||||
"state",
|
||||
}
|
||||
}
|
||||
|
||||
return cls(
|
||||
prompt=prompt,
|
||||
prompt_index=result.get("prompt_index"),
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
trajectory=result.get("trajectory"),
|
||||
trajectory_timesteps=result.get("trajectory_timesteps"),
|
||||
trajectory_decoded=result.get("trajectory_decoded"),
|
||||
video_path=result.get("video_path"),
|
||||
peak_memory_mb=result.get("peak_memory_mb"),
|
||||
state=result.get("state"),
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def to_legacy_dict(self) -> dict[str, Any]:
|
||||
result = {
|
||||
"prompts": self.prompt,
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
"trajectory": self.trajectory,
|
||||
"trajectory_timesteps": self.trajectory_timesteps,
|
||||
"trajectory_decoded": self.trajectory_decoded,
|
||||
"video_path": self.video_path,
|
||||
"peak_memory_mb": self.peak_memory_mb,
|
||||
}
|
||||
if self.prompt_index is not None:
|
||||
result["prompt_index"] = self.prompt_index
|
||||
result["prompt"] = self.prompt
|
||||
if self.state is not None:
|
||||
result["state"] = self.state
|
||||
result.update(self.extra)
|
||||
return result
|
||||
|
||||
|
||||
__all__ = ["GenerationResult"]
|
||||
@@ -0,0 +1,210 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerConfig:
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8000
|
||||
output_dir: str = "outputs/"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelismConfig:
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
hsdp_replicate_dim: int = 1
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class OffloadConfig:
|
||||
dit: bool = True
|
||||
dit_layerwise: bool = True
|
||||
text_encoder: bool = True
|
||||
image_encoder: bool = True
|
||||
vae: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompileConfig:
|
||||
enabled: bool = False
|
||||
kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QuantizationConfig:
|
||||
text_encoder_quant: str | None = None
|
||||
transformer_quant: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class EngineConfig:
|
||||
num_gpus: int = 1
|
||||
execution_backend: Literal["mp", "ray"] = "mp"
|
||||
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
|
||||
offload: OffloadConfig = field(default_factory=OffloadConfig)
|
||||
compile: CompileConfig = field(default_factory=CompileConfig)
|
||||
enable_stage_verification: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
disable_autocast: bool = False
|
||||
quantization: QuantizationConfig | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ComponentConfig:
|
||||
config_root: str | None = None
|
||||
pipeline_config_path: str | None = None
|
||||
text_encoder_weights: str | None = None
|
||||
transformer_weights: str | None = None
|
||||
transformer_2_weights: str | None = None
|
||||
vae_weights: str | None = None
|
||||
upsampler_weights: str | None = None
|
||||
lora_path: str | None = None
|
||||
override_pipeline_cls_name: str | None = None
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineSelection:
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
|
||||
profile: str | None = None
|
||||
profile_version: str | None = None
|
||||
components: ComponentConfig = field(default_factory=ComponentConfig)
|
||||
profile_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
experimental: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GeneratorConfig:
|
||||
model_path: str
|
||||
revision: str | None = None
|
||||
trust_remote_code: bool = False
|
||||
engine: EngineConfig = field(default_factory=EngineConfig)
|
||||
pipeline: PipelineSelection = field(default_factory=PipelineSelection)
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig:
|
||||
prompt_path: str | None = None
|
||||
image_path: str | list[str] | None = None
|
||||
video_path: str | list[str] | None = None
|
||||
pil_image: Any | None = None
|
||||
pose: str | None = None
|
||||
mouse_cond: Any | None = None
|
||||
keyboard_cond: Any | None = None
|
||||
grid_sizes: Any | None = None
|
||||
c2ws_plucker_emb: Any | None = None
|
||||
refine_from: str | None = None
|
||||
stage1_video: Any | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingConfig:
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestRuntimeConfig:
|
||||
enable_teacache: bool = False
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig:
|
||||
output_path: str = "outputs/"
|
||||
output_video_name: str | None = None
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
return_state: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContinuationState:
|
||||
kind: str
|
||||
payload: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlannedStage:
|
||||
name: str
|
||||
kind: str
|
||||
source: str | None = None
|
||||
overrides: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationPlan:
|
||||
stages: list[PlannedStage]
|
||||
final_stage: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationRequest:
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str | None = None
|
||||
inputs: InputConfig = field(default_factory=InputConfig)
|
||||
sampling: SamplingConfig = field(default_factory=SamplingConfig)
|
||||
runtime: RequestRuntimeConfig = field(default_factory=RequestRuntimeConfig)
|
||||
output: OutputConfig = field(default_factory=OutputConfig)
|
||||
stage_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
state: ContinuationState | None = None
|
||||
plan: GenerationPlan | None = None
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RunConfig:
|
||||
generator: GeneratorConfig
|
||||
request: GenerationRequest
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServeConfig:
|
||||
generator: GeneratorConfig
|
||||
server: ServerConfig = field(default_factory=ServerConfig)
|
||||
default_request: GenerationRequest = field(default_factory=GenerationRequest)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"ComponentConfig",
|
||||
"ContinuationState",
|
||||
"EngineConfig",
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
]
|
||||
@@ -0,0 +1,738 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Bidirectional Sparse Attention (BSA) backend for FastVideo.
|
||||
|
||||
Pure-PyTorch reference implementation from:
|
||||
"Bidirectional Sparse Attention for Faster Video Diffusion Training"
|
||||
(arXiv:2509.01085)
|
||||
|
||||
BSA sparsifies both queries (pruning redundant tokens per block) and
|
||||
key-value pairs (keeping only relevant KV blocks per query block).
|
||||
|
||||
This is a training-free inference backend: it works with any model
|
||||
trained with full attention by applying BSA sparsity at inference time.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_varlen_func_impl, )
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func as flash_attn_varlen_func_impl
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
except ImportError:
|
||||
FLASH_ATTN_AVAILABLE = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BSA_TILE_SIZE = (4, 4, 4)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cached index helpers (same pattern as VSA)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Map raster-order tokens to tile-contiguous order."""
|
||||
T, H, W = dit_seq_shape
|
||||
ts, hs, ws = tile_size
|
||||
indices = torch.arange(T * H * W, device=device, dtype=torch.long).reshape(T, H, W)
|
||||
ls = []
|
||||
for t in range(math.ceil(T / ts)):
|
||||
for h in range(math.ceil(H / hs)):
|
||||
for w in range(math.ceil(W / ws)):
|
||||
ls.append(indices[
|
||||
t * ts:min(t * ts + ts, T),
|
||||
h * hs:min(h * hs + hs, H),
|
||||
w * ws:min(w * ws + ws, W),
|
||||
].flatten())
|
||||
return torch.cat(ls, dim=0)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_reverse_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Inverse mapping: tile-contiguous order back to raster order."""
|
||||
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BSA core operations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _prune_queries(
|
||||
q_blocks: torch.Tensor,
|
||||
keep_ratio: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""
|
||||
Prune redundant query tokens within each block.
|
||||
|
||||
Scores tokens by cosine similarity to the block center.
|
||||
Keeps the LEAST similar (most informative) tokens.
|
||||
|
||||
Args:
|
||||
q_blocks: [B, N_heads, N_blocks, block_size, D]
|
||||
keep_ratio: fraction of tokens to keep
|
||||
|
||||
Returns:
|
||||
sparse_q: [B, N_heads, N_blocks, keep_size, D]
|
||||
keep_indices: [B, N_heads, N_blocks, keep_size]
|
||||
keep_size: int
|
||||
"""
|
||||
B, H, N, S, D = q_blocks.shape
|
||||
keep_size = max(1, int(S * keep_ratio))
|
||||
|
||||
if keep_size >= S:
|
||||
idx = torch.arange(S, device=q_blocks.device)
|
||||
idx = idx.view(1, 1, 1, S).expand(B, H, N, S)
|
||||
return q_blocks, idx, S
|
||||
|
||||
center_idx = S // 2
|
||||
center = q_blocks[:, :, :, center_idx:center_idx + 1, :]
|
||||
|
||||
q_norm = F.normalize(q_blocks, dim=-1)
|
||||
c_norm = F.normalize(center, dim=-1)
|
||||
similarity = (q_norm * c_norm).sum(dim=-1) # [B, H, N, S]
|
||||
|
||||
# lowest similarity = most distinctive = keep
|
||||
_, indices = similarity.topk(keep_size, dim=-1, largest=False)
|
||||
indices, _ = indices.sort(dim=-1)
|
||||
|
||||
idx_expand = indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
sparse_q = torch.gather(q_blocks, 3, idx_expand)
|
||||
|
||||
return sparse_q, indices, keep_size
|
||||
|
||||
|
||||
def _select_kv_blocks(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
cumulative_threshold: float,
|
||||
min_kv_blocks: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Dynamically select KV blocks for each query block.
|
||||
|
||||
Mean-pools to block level, computes block attention scores,
|
||||
admits blocks in descending order until cumulative mass
|
||||
exceeds threshold.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
cumulative_threshold: e.g. 0.9
|
||||
min_kv_blocks: minimum blocks to keep
|
||||
|
||||
Returns:
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
"""
|
||||
B, H, N, _, D = sparse_q.shape
|
||||
|
||||
q_repr = sparse_q.mean(dim=3)
|
||||
k_repr = k_blocks.mean(dim=3)
|
||||
|
||||
scores = torch.matmul(q_repr, k_repr.transpose(-1, -2)) / (D**0.5)
|
||||
block_attn = F.softmax(scores, dim=-1)
|
||||
|
||||
sorted_attn, sorted_idx = block_attn.sort(dim=-1, descending=True)
|
||||
cumsum = sorted_attn.cumsum(dim=-1)
|
||||
|
||||
keep_sorted = torch.ones_like(cumsum, dtype=torch.bool)
|
||||
keep_sorted[..., 1:] = cumsum[..., :-1] < cumulative_threshold
|
||||
|
||||
min_mask = torch.zeros_like(keep_sorted)
|
||||
min_mask[..., :min(min_kv_blocks, N)] = True
|
||||
keep_sorted = keep_sorted | min_mask
|
||||
|
||||
kv_mask = torch.zeros_like(block_attn, dtype=torch.bool)
|
||||
kv_mask.scatter_(-1, sorted_idx, keep_sorted)
|
||||
|
||||
return kv_mask
|
||||
|
||||
|
||||
def _compute_sparse_attention(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute attention for each query block against selected KV blocks.
|
||||
|
||||
Handles per-batch and per-head KV masks correctly.
|
||||
Uses flash_attn_varlen_func when available on GPU.
|
||||
Falls back to pure-PyTorch reference on CPU.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean (per-batch, per-head)
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
if FLASH_ATTN_AVAILABLE and sparse_q.is_cuda:
|
||||
return _compute_sparse_attention_flash(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
else:
|
||||
return _compute_sparse_attention_reference(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
|
||||
def _compute_sparse_attention_reference(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Pure-PyTorch fallback with per-batch, per-head mask support."""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for qb in range(N):
|
||||
selected = kv_mask[b, h, qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
# [num_sel * Sk, D]
|
||||
sel_k = k_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
sel_v = v_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
|
||||
q = sparse_q[b, h, qb] # [Sq, D]
|
||||
scores = torch.matmul(q, sel_k.transpose(-1, -2)) / (D**0.5)
|
||||
weights = F.softmax(scores, dim=-1)
|
||||
output[b, h, qb] = torch.matmul(weights, sel_v)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _compute_sparse_attention_flash(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
FlashAttention implementation with per-batch, per-head mask support.
|
||||
|
||||
Strategy: check if all heads share the same mask. If so, use a single
|
||||
FlashAttention call per batch (fast path). If not, process each head
|
||||
separately (correct path).
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
Sk = k_blocks.shape[3]
|
||||
device = sparse_q.device
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
# Check if all heads share the same mask for this batch element
|
||||
# Compare each head's mask to head 0's mask
|
||||
head0_mask = kv_mask[b, 0] # [N, N]
|
||||
all_heads_same = all(torch.equal(kv_mask[b, h], head0_mask) for h in range(1, H))
|
||||
|
||||
if all_heads_same:
|
||||
# Fast path: all heads share the same mask, single FA call
|
||||
_flash_attn_single_mask(
|
||||
sparse_q[b],
|
||||
k_blocks[b],
|
||||
v_blocks[b],
|
||||
head0_mask,
|
||||
output[b],
|
||||
H,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
else:
|
||||
# Per-head path: process each head individually
|
||||
for h in range(H):
|
||||
head_mask = kv_mask[b, h] # [N, N]
|
||||
# Process single head: squeeze head dim, run FA, put back
|
||||
_flash_attn_single_head(
|
||||
sparse_q[b, h],
|
||||
k_blocks[b, h],
|
||||
v_blocks[b, h],
|
||||
head_mask,
|
||||
output,
|
||||
b,
|
||||
h,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _flash_attn_single_mask(
|
||||
sparse_q_b: torch.Tensor, # [H, N, Sq, D]
|
||||
k_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
v_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output_b: torch.Tensor, # [H, N, Sq, D] (modified in-place)
|
||||
H: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for all heads sharing the same KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [H, Sq, D] -> [Sq, H, D]
|
||||
q_block = sparse_q_b[:, qb].permute(1, 0, 2)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [H, num_sel, Sk, D] -> [num_kv_tokens, H, D]
|
||||
sel_k = k_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
sel_v = v_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = int((cu_seqlens_k_t[1:] - cu_seqlens_k_t[:-1]).max().item())
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, H, D]
|
||||
output_b[:, qb] = block_out.permute(1, 0, 2) # [H, Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _flash_attn_single_head(
|
||||
sparse_q_bh: torch.Tensor, # [N, Sq, D]
|
||||
k_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
v_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output: torch.Tensor, # [B, H, N, Sq, D] (modified in-place)
|
||||
b: int,
|
||||
h: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for a single head with its own KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb]
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [Sq, D] -> [Sq, 1, D] (single head)
|
||||
q_block = sparse_q_bh[qb].unsqueeze(1)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [num_sel, Sk, D] -> [num_kv_tokens, 1, D]
|
||||
sel_k = k_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
sel_v = v_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = int((cu_seqlens_k_t[1:] - cu_seqlens_k_t[:-1]).max().item())
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, 1, D]
|
||||
output[b, h, qb] = block_out.squeeze(1) # [Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _reconstruct_pruned(
|
||||
sparse_output: torch.Tensor,
|
||||
keep_indices: torch.Tensor,
|
||||
block_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Scatter sparse output back to full block size.
|
||||
Pruned positions get nearest kept token's output.
|
||||
|
||||
Handles per-batch, per-head indices correctly.
|
||||
|
||||
Args:
|
||||
sparse_output: [B, H, N, keep_size, D]
|
||||
keep_indices: [B, H, N, keep_size]
|
||||
block_size: original tokens per block
|
||||
|
||||
Returns:
|
||||
full_output: [B, H, N, block_size, D]
|
||||
"""
|
||||
B, H, N, keep_size, D = sparse_output.shape
|
||||
device = sparse_output.device
|
||||
|
||||
if keep_size >= block_size:
|
||||
return sparse_output
|
||||
|
||||
full_output = torch.zeros(B, H, N, block_size, D, device=device, dtype=sparse_output.dtype)
|
||||
|
||||
# Scatter kept tokens
|
||||
idx_expand = keep_indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
full_output.scatter_(3, idx_expand, sparse_output)
|
||||
|
||||
# Fill pruned positions with nearest kept token (vectorized)
|
||||
all_pos = torch.arange(block_size, device=device)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for n in range(N):
|
||||
kept = keep_indices[b, h, n] # [keep_size]
|
||||
|
||||
# Distance from every position to every kept position
|
||||
dists = (all_pos.view(-1, 1) - kept.view(1, -1)).abs()
|
||||
nearest_local_idx = dists.argmin(dim=1) # [block_size]
|
||||
|
||||
# Identify pruned positions
|
||||
is_pruned = torch.ones(block_size, dtype=torch.bool, device=device)
|
||||
is_pruned[kept] = False
|
||||
pruned_indices = is_pruned.nonzero(as_tuple=True)[0]
|
||||
|
||||
if pruned_indices.numel() > 0:
|
||||
src_indices = nearest_local_idx[pruned_indices]
|
||||
full_output[b, h, n, pruned_indices] = sparse_output[b, h, n, src_indices]
|
||||
|
||||
return full_output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastVideo backend classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BSAAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = False
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "BSA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["BSAAttentionImpl"]:
|
||||
return BSAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["BSAAttentionMetadata"]:
|
||||
return BSAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["BSAAttentionMetadataBuilder"]:
|
||||
return BSAAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class BSAAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
dit_seq_shape: tuple[int, int, int]
|
||||
total_seq_length: int
|
||||
num_blocks: int
|
||||
block_size: int
|
||||
tile_partition_indices: torch.LongTensor
|
||||
reverse_tile_partition_indices: torch.LongTensor
|
||||
# BSA-specific config
|
||||
query_keep_ratio: float
|
||||
kv_cumulative_threshold: float
|
||||
min_kv_blocks: int
|
||||
|
||||
|
||||
class BSAAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
bsa_query_keep_ratio: float = 0.5,
|
||||
bsa_kv_cumulative_threshold: float = 0.9,
|
||||
bsa_min_kv_blocks: int = 4,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> "BSAAttentionMetadata":
|
||||
# Ensure patching does not drop tokens silently.
|
||||
assert all(r % p == 0 for r, p in zip(raw_latent_shape, patch_size, strict=False)), (
|
||||
"raw_latent_shape must be divisible by patch_size for BSA", )
|
||||
|
||||
dit_seq_shape = (
|
||||
raw_latent_shape[0] // patch_size[0],
|
||||
raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2],
|
||||
)
|
||||
|
||||
total_seq_length = math.prod(dit_seq_shape)
|
||||
block_size = math.prod(BSA_TILE_SIZE)
|
||||
# Require exact tiling to avoid reshape failures later.
|
||||
assert all(d % t == 0 for d, t in zip(dit_seq_shape, BSA_TILE_SIZE, strict=False)), (
|
||||
"dit_seq_shape must be divisible by BSA_TILE_SIZE", )
|
||||
num_blocks = total_seq_length // block_size
|
||||
|
||||
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
|
||||
return BSAAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
total_seq_length=total_seq_length,
|
||||
num_blocks=num_blocks,
|
||||
block_size=block_size,
|
||||
tile_partition_indices=tile_partition_indices,
|
||||
reverse_tile_partition_indices=reverse_tile_partition_indices,
|
||||
query_keep_ratio=bsa_query_keep_ratio,
|
||||
kv_cumulative_threshold=bsa_kv_cumulative_threshold,
|
||||
min_kv_blocks=bsa_min_kv_blocks,
|
||||
)
|
||||
|
||||
|
||||
class BSAAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.prefix = prefix
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
if num_kv_heads is not None and num_kv_heads != num_heads:
|
||||
raise ValueError("BSA backend does not support grouped-query attention")
|
||||
if causal:
|
||||
raise ValueError("BSA backend is bidirectional; causal=True is unsupported")
|
||||
if softmax_scale is not None:
|
||||
expected_scale = 1.0 / math.sqrt(self.head_size)
|
||||
if not math.isclose(softmax_scale, expected_scale, rel_tol=1e-4, abs_tol=1e-5):
|
||||
raise ValueError("softmax_scale must be default (1/sqrt(d)) for BSA")
|
||||
try:
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
except (AssertionError, RuntimeError):
|
||||
self.sp_size = 1
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from raster order to tile-contiguous order."""
|
||||
# qkv: [B, L, num_heads, D]
|
||||
return qkv[:, attn_metadata.tile_partition_indices]
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from tile-contiguous order back to raster order."""
|
||||
return output[:, attn_metadata.reverse_tile_partition_indices]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
BSA attention forward pass.
|
||||
|
||||
Input tensors are already in tile-contiguous order from preprocess_qkv.
|
||||
|
||||
Args:
|
||||
query: [B, L, num_heads, D] (tile-ordered)
|
||||
key: [B, L, num_heads, D] (tile-ordered)
|
||||
value: [B, L, num_heads, D] (tile-ordered)
|
||||
attn_metadata: BSA metadata
|
||||
|
||||
Returns:
|
||||
output: [B, L, num_heads, D] (tile-ordered)
|
||||
"""
|
||||
B, L, H, D = query.shape
|
||||
block_size = attn_metadata.block_size
|
||||
num_blocks = attn_metadata.num_blocks
|
||||
assert num_blocks * block_size == L, "Sequence length must match tiling"
|
||||
|
||||
# Reshape to [B, H, L, D] for attention computation
|
||||
q = query.transpose(1, 2).contiguous() # [B, H, L, D]
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Reshape into blocks: [B, H, num_blocks, block_size, D]
|
||||
q_blocks = q.view(B, H, num_blocks, block_size, D)
|
||||
k_blocks = k.view(B, H, num_blocks, block_size, D)
|
||||
v_blocks = v.view(B, H, num_blocks, block_size, D)
|
||||
|
||||
# --- Query sparsification ---
|
||||
sparse_q, keep_indices, keep_size = _prune_queries(q_blocks, attn_metadata.query_keep_ratio)
|
||||
|
||||
# --- KV block selection ---
|
||||
kv_mask = _select_kv_blocks(
|
||||
sparse_q,
|
||||
k_blocks,
|
||||
attn_metadata.kv_cumulative_threshold,
|
||||
attn_metadata.min_kv_blocks,
|
||||
)
|
||||
|
||||
# --- Sparse attention ---
|
||||
sparse_output = _compute_sparse_attention(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
# --- Reconstruct pruned positions ---
|
||||
full_output = _reconstruct_pruned(sparse_output, keep_indices, block_size)
|
||||
|
||||
# Reshape back: [B, H, num_blocks, block_size, D] -> [B, H, L, D] -> [B, L, H, D]
|
||||
hidden_states = full_output.view(B, H, L, D).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,188 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CArchConfig(DiTArchConfig):
|
||||
"""Configuration for GEN3C architecture (VideoExtendGeneralDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Official GEN3C checkpoint key naming to FastVideo mapping.
|
||||
# The official checkpoint uses nn.Sequential patterns like attn.to_q.0 (Linear)
|
||||
# and attn.to_q.1 (RMSNorm), and layer1/layer2 for MLP.
|
||||
#
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_*.weight -> time_embed.t_embedder.linear_*.weight
|
||||
r"^net\.t_embedder\.0\.(.*)$": r"time_embed.time_proj.\1",
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$": r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$": r"time_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Augment sigma embedding (GEN3C-specific)
|
||||
r"^net\.augment_sigma_embedder\.0\.(.*)$": r"augment_sigma_embed.time_proj.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_1\.(.*)$": r"augment_sigma_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_2\.(.*)$": r"augment_sigma_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Affine embedding norm: net.affline_norm.weight -> affine_norm.weight
|
||||
# Note: "affline" is a typo in the official GEN3C checkpoint (should be "affine")
|
||||
r"^net\.affline_norm\.(.*)$": r"affine_norm.\1",
|
||||
|
||||
# Extra positional embeddings (learnable per-axis)
|
||||
r"^net\.extra_pos_embedder\.pos_emb_t$": r"learnable_pos_embed.pos_emb_t",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_h$": r"learnable_pos_embed.pos_emb_h",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_w$": r"learnable_pos_embed.pos_emb_w",
|
||||
|
||||
# Transformer blocks: net.blocks.blockN -> transformer_blocks.N
|
||||
# Official uses: block.attn.to_q.0 (Linear), block.attn.to_q.1 (QK RMSNorm)
|
||||
#
|
||||
# Self-attention (block index 0)
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
# AdaLN modulation for self-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.\2",
|
||||
|
||||
# Cross-attention (block index 1)
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
# AdaLN modulation for cross-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.\2",
|
||||
|
||||
# MLP (block index 2): layer1 -> fc_in, layer2 -> fc_out
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
# AdaLN modulation for MLP
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.\2",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN: net.final_layer.adaLN_modulation -> final_layer.adaln_modulation
|
||||
r"^net\.final_layer\.adaLN_modulation\.(.*)$": r"final_layer.adaln_modulation.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (rope position embeddings computed dynamically)
|
||||
# - net.accum_* keys (training metadata)
|
||||
# - logvar.* (training-only module, not used in inference)
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$": r"transformer_blocks.\1.mlp.\2",
|
||||
})
|
||||
|
||||
# GEN3C architecture parameters
|
||||
# Base VAE latent channels
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
|
||||
# Channels per 3D cache buffer: 16 (warped frame latent) + 16 (warped mask latent)
|
||||
CHANNELS_PER_BUFFER: int = 32
|
||||
|
||||
# Number of 3D cache buffers
|
||||
frame_buffer_max: int = 2
|
||||
|
||||
# Attention configuration (7B model: 32 heads x 128 dim = 4096 hidden)
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128 # 4096 / 32
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
|
||||
# Text encoder configuration
|
||||
text_embed_dim: int = 1024
|
||||
|
||||
# AdaLN-LoRA configuration
|
||||
adaln_lora_dim: int = 256
|
||||
use_adaln_lora: bool = True
|
||||
|
||||
# GEN3C-specific: augment sigma embedding for conditioning noise augmentation
|
||||
# Note: The official GEN3C-Cosmos-7B checkpoint was trained without this
|
||||
add_augment_sigma_embedding: bool = False
|
||||
|
||||
# Position embedding configuration
|
||||
max_size: tuple[int, int, int] = (128, 240, 240) # T, H, W
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0) # T, H, W scaling
|
||||
|
||||
# GEN3C uses learnable positional embeddings in addition to RoPE
|
||||
extra_pos_embed_type: str = "learnable"
|
||||
|
||||
# Padding mask handling
|
||||
concat_padding_mask: bool = True
|
||||
|
||||
# Cross-attention projection (not used in GEN3C 7B)
|
||||
use_crossattn_projection: bool = False
|
||||
|
||||
# RoPE FPS modulation
|
||||
rope_enable_fps_modulation: bool = True
|
||||
|
||||
# QK normalization
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
|
||||
# Affine embedding normalization
|
||||
affine_emb_norm: bool = True
|
||||
|
||||
# Block format (THWBD for GEN3C compatibility)
|
||||
block_x_format: str = "THWBD"
|
||||
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
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.in_channels
|
||||
|
||||
# Calculate total input channels for patch embedding:
|
||||
# - in_channels (16): VAE latent
|
||||
# - condition_video_input_mask (1): Binary mask for conditioning frames
|
||||
# - condition_video_pose (frame_buffer_max * 32): 3D cache buffers
|
||||
# - padding_mask (1 if concat_padding_mask): Padding mask
|
||||
self.buffer_channels = self.frame_buffer_max * self.CHANNELS_PER_BUFFER
|
||||
self.total_input_channels = (
|
||||
self.in_channels + # 16: VAE latent
|
||||
1 + # 1: condition_video_input_mask
|
||||
self.buffer_channels # 64: 3D cache buffers (2 * 32)
|
||||
)
|
||||
# padding_mask is added in build_patch_embed if concat_padding_mask=True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CVideoConfig(DiTConfig):
|
||||
"""Configuration for GEN3C video generation model."""
|
||||
arch_config: DiTArchConfig = field(default_factory=Gen3CArchConfig)
|
||||
prefix: str = "Gen3C"
|
||||
@@ -1,6 +1,7 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -12,6 +13,7 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CVAEConfig(CosmosVAEConfig):
|
||||
"""
|
||||
GEN3C VAE config placeholder.
|
||||
|
||||
GEN3C uses tokenizer-backed VAE loading logic at runtime, but we keep a
|
||||
model-specific config class so pipeline/model configs stay model-scoped.
|
||||
"""
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
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.gen3c import Gen3CVideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import (T5LargeArchConfig, T5LargeConfig)
|
||||
from fastvideo.configs.models.vaes import Gen3CVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Gen3CT5LargeArchConfig(T5LargeArchConfig):
|
||||
"""T5 Large arch config that pads inputs to max_length.
|
||||
|
||||
GEN3C requires padded text encoder inputs, while the base
|
||||
T5 config no longer pads by default after the SP mask
|
||||
refactor [PR#1142](https://github.com/hao-ai-lab/FastVideo/pull/1142).
|
||||
"""
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Gen3CT5LargeConfig(T5LargeConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=_Gen3CT5LargeArchConfig)
|
||||
prefix: str = "t5"
|
||||
|
||||
|
||||
def t5_large_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Postprocess T5 Large text encoder outputs for GEN3C pipeline.
|
||||
|
||||
Return raw last_hidden_state without truncation/padding.
|
||||
"""
|
||||
hidden_state = outputs.last_hidden_state
|
||||
|
||||
if hidden_state is None:
|
||||
raise ValueError("T5 Large outputs missing last_hidden_state")
|
||||
|
||||
nan_count = torch.isnan(hidden_state).sum()
|
||||
if nan_count > 0:
|
||||
hidden_state = hidden_state.masked_fill(torch.isnan(hidden_state), 0.0)
|
||||
|
||||
# Zero out embeddings beyond actual sequence length (vectorized)
|
||||
if outputs.attention_mask is not None:
|
||||
attention_mask = outputs.attention_mask
|
||||
lengths = attention_mask.sum(dim=1)
|
||||
max_len = hidden_state.shape[1]
|
||||
mask = torch.arange(max_len, device=hidden_state.device)[None, :] >= lengths[:, None]
|
||||
hidden_state[mask] = 0.0
|
||||
|
||||
return hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CConfig(PipelineConfig):
|
||||
"""Configuration for GEN3C Video Generation Pipeline.
|
||||
|
||||
GEN3C extends Cosmos with 3D cache for camera-controlled video generation.
|
||||
Key parameters:
|
||||
- frame_buffer_max: Number of 3D cache buffers (default: 2)
|
||||
- noise_aug_strength: Strength of noise augmentation per buffer
|
||||
- filter_points_threshold: Threshold for filtering unreliable depth points
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Gen3CVideoConfig)
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=Gen3CVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (_Gen3CT5LargeConfig(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5_large_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# GEN3C-specific conditioning parameters
|
||||
conditioning_strategy: str = "frame_replace"
|
||||
min_num_conditional_frames: int = 1
|
||||
max_num_conditional_frames: int = 2
|
||||
# Match official GEN3C/Cosmos inference defaults.
|
||||
sigma_conditional: float = 0.001
|
||||
sigma_data: float = 0.5
|
||||
state_ch: int = 16
|
||||
state_t: int = 16 # GEN3C uses 16 latent frames (121 pixel frames)
|
||||
text_encoder_class: str = "T5"
|
||||
|
||||
# Flow matching parameters
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: float = 1.0
|
||||
|
||||
# GEN3C 3D Cache parameters
|
||||
frame_buffer_max: int = 2
|
||||
noise_aug_strength: float = 0.0
|
||||
filter_points_threshold: float = 0.05
|
||||
|
||||
# Depth estimation settings
|
||||
use_moge_depth: bool = True
|
||||
moge_model_name: str = "Ruicheng/moge-vitl"
|
||||
offload_moge_after_depth: bool = True
|
||||
|
||||
# Camera trajectory settings (matching NVIDIA inference defaults)
|
||||
default_trajectory_type: str = "left"
|
||||
default_movement_distance: float = 0.3
|
||||
default_camera_rotation: str = "center_facing"
|
||||
|
||||
# Video generation settings
|
||||
# Match official GEN3C defaults (height=704, width=1280).
|
||||
video_resolution: tuple[int, int] = (704, 1280) # H, W
|
||||
num_frames: int = 121 # Default number of frames to generate
|
||||
|
||||
# Generation frame rate
|
||||
fps: int = 24
|
||||
|
||||
# Explicit CFG behavior policy:
|
||||
# - "legacy": CFG branch only when guidance_scale > 1.0
|
||||
# - "official_uncond_at_unity": also run uncond branch at guidance_scale == 1.0
|
||||
cfg_behavior: str = "legacy"
|
||||
default_negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality.")
|
||||
|
||||
# Autoregressive generation settings
|
||||
autoregressive_chunk_frames: int = 121 # Frames per chunk
|
||||
autoregressive_overlap_frames: int = 1 # Overlap between chunks
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
self._vae_latent_dim = 16
|
||||
|
||||
# Validate frame buffer configuration matches DiT
|
||||
if hasattr(self.dit_config, 'arch_config'):
|
||||
arch_config = self.dit_config.arch_config
|
||||
if (hasattr(arch_config, 'frame_buffer_max') and arch_config.frame_buffer_max != self.frame_buffer_max):
|
||||
raise ValueError(f"frame_buffer_max mismatch: pipeline config has {self.frame_buffer_max}, "
|
||||
f"DiT config has {arch_config.frame_buffer_max}")
|
||||
|
||||
allowed_cfg_behavior = {"legacy", "official_uncond_at_unity"}
|
||||
if self.cfg_behavior not in allowed_cfg_behavior:
|
||||
raise ValueError(f"cfg_behavior must be one of {sorted(allowed_cfg_behavior)}, got {self.cfg_behavior!r}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CInferenceConfig(Gen3CConfig):
|
||||
"""Configuration for GEN3C inference with optimized defaults."""
|
||||
|
||||
# Use smaller batch sizes for inference
|
||||
batch_size: int = 1
|
||||
|
||||
# Enable gradient checkpointing for memory efficiency
|
||||
gradient_checkpointing: bool = False
|
||||
|
||||
# Inference-specific parameters
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 35
|
||||
|
||||
# Disable noise augmentation during inference
|
||||
noise_aug_strength: float = 0.0
|
||||
@@ -72,6 +72,14 @@ class SamplingParam:
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
|
||||
# GEN3C camera control
|
||||
trajectory_type: str | None = None
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
@@ -96,17 +104,42 @@ class SamplingParam:
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.registry import get_sampling_param_cls_for_name
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
from fastvideo.registry import _get_config_info
|
||||
config_info = _get_config_info(
|
||||
model_path,
|
||||
raise_on_missing=False,
|
||||
)
|
||||
|
||||
if config_info is not None and config_info.default_profile:
|
||||
return cls._from_profile(config_info.default_profile)
|
||||
|
||||
sampling_cls = (config_info.sampling_param_cls if config_info is not None else None)
|
||||
|
||||
if sampling_cls is not None:
|
||||
sampling_param: SamplingParam = sampling_cls()
|
||||
else:
|
||||
logger.warning("Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param "
|
||||
"for %s. Using the default sampling param.",
|
||||
model_path,
|
||||
)
|
||||
sampling_param = cls()
|
||||
|
||||
return sampling_param
|
||||
|
||||
@classmethod
|
||||
def _from_profile(cls, profile_name: str) -> "SamplingParam":
|
||||
from fastvideo.configs.sample.profiles import get_profile
|
||||
profile = get_profile(profile_name)
|
||||
if profile is None:
|
||||
raise ValueError(f"Profile {profile_name!r} not found in "
|
||||
"profile registry")
|
||||
instance = cls()
|
||||
for key, value in profile.defaults.items():
|
||||
setattr(instance, key, value)
|
||||
instance.__post_init__()
|
||||
return instance
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: Any) -> Any:
|
||||
"""Add CLI arguments for SamplingParam fields"""
|
||||
|
||||
@@ -1,18 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_2B_Video2World_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 93
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
|
||||
num_inference_steps: int = 35
|
||||
# Migrated to profile-based defaults.
|
||||
# See fastvideo/pipelines/basic/cosmos/profiles.py
|
||||
|
||||
@@ -1,23 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25SamplingParamBase(SamplingParam):
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 77
|
||||
fps: int = 24
|
||||
seed: int = 0
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
|
||||
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
|
||||
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
num_inference_steps: int = 35
|
||||
# Migrated to profile-based defaults.
|
||||
# See fastvideo/pipelines/basic/cosmos/profiles.py
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Migrated to profile-based defaults.
|
||||
# See fastvideo/pipelines/basic/gen3c/profiles.py
|
||||
@@ -0,0 +1,33 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Profile-based defaults for SamplingParam.
|
||||
|
||||
A ModelProfile captures the recommended sampling defaults for a specific
|
||||
model variant (resolution, fps, guidance scale, etc.) without requiring
|
||||
a dedicated SamplingParam subclass. ``SamplingParam.from_pretrained``
|
||||
resolves the profile via the registry and applies its ``defaults`` dict
|
||||
with simple ``setattr`` calls on a base ``SamplingParam`` instance.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
# Global registry: profile name -> ModelProfile
|
||||
_PROFILE_REGISTRY: dict[str, ModelProfile] = {}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelProfile:
|
||||
"""Declarative bag of sampling defaults for one model variant."""
|
||||
|
||||
name: str
|
||||
defaults: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_PROFILE_REGISTRY[self.name] = self
|
||||
|
||||
|
||||
def get_profile(name: str) -> ModelProfile | None:
|
||||
"""Look up a registered profile by name."""
|
||||
return _PROFILE_REGISTRY.get(name)
|
||||
@@ -7,7 +7,7 @@ Example usage:
|
||||
# launch a server and benchmark on it
|
||||
|
||||
# T2V or T2I or any other multimodal generation model
|
||||
fastvideo serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers --port 8000
|
||||
fastvideo serve --config serve.yaml
|
||||
|
||||
# benchmark it and make sure the port is the same as the server's port
|
||||
fastvideo bench --dataset vbench --num-prompts 20 --port 8000
|
||||
|
||||
@@ -2,19 +2,17 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.utils import RaiseNotImplementedAction
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_generate_run_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_RUN_CONFIG_ATTR = "_fastvideo_validated_run_config"
|
||||
|
||||
|
||||
class GenerateSubcommand(CLISubcommand):
|
||||
@@ -23,89 +21,47 @@ class GenerateSubcommand(CLISubcommand):
|
||||
def __init__(self) -> None:
|
||||
self.name = "generate"
|
||||
super().__init__()
|
||||
self.init_arg_names = self._get_init_arg_names()
|
||||
self.generation_arg_names = self._get_generation_arg_names()
|
||||
|
||||
def _get_init_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for VideoGenerator initialization"""
|
||||
return ["num_gpus", "tp_size", "sp_size", "model_path"]
|
||||
|
||||
def _get_generation_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for generate_video method"""
|
||||
return [field.name for field in dataclasses.fields(SamplingParam)]
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = ['subparser', 'config', 'dispatch_function']
|
||||
run_config = getattr(args, _VALIDATED_RUN_CONFIG_ATTR, None)
|
||||
if run_config is None:
|
||||
run_config = build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
logger.info("CLI generate config: %s", run_config)
|
||||
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (k not in excluded_args and v is not None and hasattr(args, '_provided') and k in args._provided):
|
||||
provided_args[k] = v
|
||||
|
||||
if 'model_path' in vars(args) and args.model_path is not None:
|
||||
provided_args['model_path'] = args.model_path
|
||||
|
||||
if 'prompt' in vars(args) and args.prompt is not None:
|
||||
provided_args['prompt'] = args.prompt
|
||||
|
||||
merged_args = {**provided_args}
|
||||
|
||||
logger.info('CLI Args: %s', merged_args)
|
||||
|
||||
if 'model_path' not in merged_args or not merged_args['model_path']:
|
||||
raise ValueError("model_path must be provided either in config file or via --model-path")
|
||||
|
||||
# Check if either prompt or prompt_txt is provided
|
||||
has_prompt = 'prompt' in merged_args and merged_args['prompt']
|
||||
has_prompt_txt = 'prompt_txt' in merged_args and merged_args['prompt_txt']
|
||||
|
||||
if not (has_prompt or has_prompt_txt):
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
|
||||
if has_prompt and has_prompt_txt:
|
||||
raise ValueError("Cannot provide both 'prompt' and 'prompt_txt'. Use only one of them.")
|
||||
|
||||
init_args = {k: v for k, v in merged_args.items() if k not in self.generation_arg_names}
|
||||
generation_args = {k: v for k, v in merged_args.items() if k in self.generation_arg_names}
|
||||
generation_args.setdefault("return_frames", False)
|
||||
|
||||
model_path = init_args.pop('model_path')
|
||||
prompt = generation_args.pop('prompt', None)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path=model_path, **init_args)
|
||||
|
||||
# Call generate_video - it handles both single and batch modes
|
||||
generator.generate_video(prompt=prompt, **generation_args)
|
||||
generator = VideoGenerator.from_config(run_config.generator)
|
||||
generator.generate(run_config.request)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
"""Validate the arguments for this command"""
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
if args.config and not os.path.exists(args.config):
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo generate requires --config PATH; use a nested "
|
||||
"run config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_RUN_CONFIG_ATTR,
|
||||
build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
generate_parser = subparsers.add_parser(
|
||||
"generate",
|
||||
help="Run inference on a model",
|
||||
usage="fastvideo generate (--model-path MODEL_PATH_OR_ID --prompt PROMPT) | --config CONFIG_FILE [OPTIONS]")
|
||||
usage="fastvideo generate --config RUN_CONFIG [--dotted.override VALUE]")
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default='',
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional."
|
||||
)
|
||||
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
generate_parser = SamplingParam.add_cli_args(generate_parser)
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--text-encoder-configs",
|
||||
action=RaiseNotImplementedAction,
|
||||
help="JSON array of text encoder configurations (NOT YET IMPLEMENTED)",
|
||||
help="Path to a nested run config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.parser import load_raw_config, parse_config
|
||||
from fastvideo.api.schema import RunConfig, ServeConfig
|
||||
|
||||
_GENERATE_OVERRIDE_PREFIXES = ("generator.", "request.")
|
||||
_SERVE_OVERRIDE_PREFIXES = (
|
||||
"generator.",
|
||||
"server.",
|
||||
"default_request.",
|
||||
)
|
||||
|
||||
|
||||
def build_generate_run_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> RunConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_GENERATE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
_ensure_generate_cli_defaults(raw)
|
||||
config = parse_config(RunConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
_validate_generate_prompt_sources(config)
|
||||
return config
|
||||
|
||||
|
||||
def build_serve_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> ServeConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("server", {})
|
||||
raw.setdefault("default_request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_SERVE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
config = parse_config(ServeConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
return config
|
||||
|
||||
|
||||
def _load_nested_config(path: str | None) -> dict[str, Any]:
|
||||
if not path:
|
||||
raise ValueError("Inference CLI requires --config PATH; use a nested config file "
|
||||
"plus optional dotted overrides")
|
||||
|
||||
raw = load_raw_config(path)
|
||||
if not isinstance(raw.get("generator"), Mapping):
|
||||
raise ValueError("Inference config must use the nested schema with a top-level "
|
||||
"'generator' mapping")
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
|
||||
def _apply_dotted_overrides(
|
||||
raw: Mapping[str, Any],
|
||||
overrides: list[str] | None,
|
||||
*,
|
||||
allowed_prefixes: tuple[str, ...],
|
||||
) -> dict[str, Any]:
|
||||
if not overrides:
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
parsed = parse_cli_overrides(overrides)
|
||||
for key in parsed:
|
||||
if "." not in key:
|
||||
raise ValueError("CLI overrides must use dotted config paths like "
|
||||
"--request.sampling.seed 42")
|
||||
if not key.startswith(allowed_prefixes):
|
||||
allowed = ", ".join(allowed_prefixes)
|
||||
raise ValueError(f"Unsupported override path {key!r}. Allowed prefixes: {allowed}")
|
||||
return apply_overrides(raw, parsed)
|
||||
|
||||
|
||||
def _ensure_generate_cli_defaults(raw: dict[str, Any]) -> None:
|
||||
request = raw.setdefault("request", {})
|
||||
output = request.setdefault("output", {})
|
||||
output.setdefault("return_frames", False)
|
||||
|
||||
|
||||
def _validate_generate_prompt_sources(config: RunConfig) -> None:
|
||||
has_prompt = config.request.prompt is not None
|
||||
has_prompt_path = config.request.inputs.prompt_path is not None
|
||||
if not (has_prompt or has_prompt_path):
|
||||
raise ValueError("Either request.prompt or request.inputs.prompt_path must be provided")
|
||||
if has_prompt and has_prompt_path:
|
||||
raise ValueError("Cannot provide both request.prompt and request.inputs.prompt_path")
|
||||
|
||||
|
||||
def _validate_num_gpus(num_gpus: int) -> None:
|
||||
if num_gpus <= 0:
|
||||
raise ValueError(f"generator.engine.num_gpus must be > 0; got {num_gpus}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_generate_run_config",
|
||||
"build_serve_config",
|
||||
]
|
||||
@@ -1,6 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
@@ -27,14 +26,17 @@ def main() -> None:
|
||||
for cmd in cmd_init():
|
||||
cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)
|
||||
cmds[cmd.name] = cmd
|
||||
args = parser.parse_args()
|
||||
|
||||
args, unknown = parser.parse_known_args()
|
||||
if unknown and args.subparser not in {"generate", "serve"}:
|
||||
parser.error(f"unrecognized arguments: {' '.join(unknown)}")
|
||||
args._unknown = unknown
|
||||
if args.subparser in cmds:
|
||||
cmds[args.subparser].validate(args)
|
||||
|
||||
if hasattr(args, "dispatch_function"):
|
||||
args.dispatch_function(args)
|
||||
else:
|
||||
parser.print_help()
|
||||
return
|
||||
|
||||
parser.print_help()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -2,14 +2,18 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.api.compat import generator_config_to_fastvideo_args
|
||||
from fastvideo.api.request_metadata import EXPLICIT_REQUEST_ATTR
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_serve_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_SERVE_CONFIG_ATTR = "_fastvideo_validated_serve_config"
|
||||
|
||||
|
||||
class ServeSubcommand(CLISubcommand):
|
||||
@@ -20,94 +24,67 @@ class ServeSubcommand(CLISubcommand):
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = {
|
||||
"subparser",
|
||||
"config",
|
||||
"dispatch_function",
|
||||
"host",
|
||||
"port",
|
||||
"output_dir",
|
||||
}
|
||||
|
||||
provided: set[str] = getattr(args, '_provided', set())
|
||||
cli_kwargs = {}
|
||||
for k, v in vars(args).items():
|
||||
if k in excluded_args:
|
||||
continue
|
||||
if k == '_provided':
|
||||
continue
|
||||
if k in provided and v is not None:
|
||||
cli_kwargs[k] = v
|
||||
|
||||
if 'model_path' not in cli_kwargs and args.model_path is not None:
|
||||
cli_kwargs['model_path'] = args.model_path
|
||||
|
||||
if not cli_kwargs.get('model_path'):
|
||||
raise ValueError("model_path must be provided via --model-path")
|
||||
serve_config = getattr(args, _VALIDATED_SERVE_CONFIG_ATTR, None)
|
||||
if serve_config is None:
|
||||
serve_config = build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
explicit_raw = getattr(
|
||||
serve_config.default_request,
|
||||
EXPLICIT_REQUEST_ATTR,
|
||||
None,
|
||||
)
|
||||
if explicit_raw:
|
||||
raise NotImplementedError("ServeConfig.default_request is not wired into the OpenAI "
|
||||
"server yet")
|
||||
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
run_server,
|
||||
run_server, )
|
||||
|
||||
logger.info("CLI serve config: %s", serve_config)
|
||||
logger.info(
|
||||
"Server will listen on %s:%d",
|
||||
serve_config.server.host,
|
||||
serve_config.server.port,
|
||||
)
|
||||
|
||||
host = getattr(args, "host", DEFAULT_HOST)
|
||||
port = getattr(args, "port", DEFAULT_PORT)
|
||||
output_dir = getattr(args, "output_dir", DEFAULT_OUTPUT_DIR)
|
||||
|
||||
logger.info("CLI serve args: %s", cli_kwargs)
|
||||
logger.info("Server will listen on %s:%d", host, port)
|
||||
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**cli_kwargs)
|
||||
run_server(fastvideo_args, host=host, port=port, output_dir=output_dir)
|
||||
fastvideo_args = generator_config_to_fastvideo_args(serve_config.generator)
|
||||
run_server(
|
||||
fastvideo_args,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
output_dir=serve_config.server.output_dir,
|
||||
)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo serve requires --config PATH; use a nested "
|
||||
"serve config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_SERVE_CONFIG_ATTR,
|
||||
build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
serve_parser = subparsers.add_parser(
|
||||
"serve",
|
||||
help="Start an OpenAI-compatible HTTP server",
|
||||
usage=("fastvideo serve --model-path MODEL_PATH_OR_ID "
|
||||
"[--host HOST] [--port PORT] [OPTIONS]"),
|
||||
)
|
||||
|
||||
serve_parser.add_argument(
|
||||
"--host",
|
||||
type=str,
|
||||
default=DEFAULT_HOST,
|
||||
help=f"Host to bind the server to (default: {DEFAULT_HOST})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
default=DEFAULT_PORT,
|
||||
help=f"Port to listen on (default: {DEFAULT_PORT})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--output-dir",
|
||||
type=str,
|
||||
default=DEFAULT_OUTPUT_DIR,
|
||||
help=("Directory for generated outputs "
|
||||
f"(default: {DEFAULT_OUTPUT_DIR})"),
|
||||
usage="fastvideo serve --config SERVE_CONFIG [--dotted.override VALUE]",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file.",
|
||||
help="Path to a nested config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
serve_parser = FastVideoArgs.add_cli_args(serve_parser)
|
||||
return cast(FlexibleArgumentParser, serve_parser)
|
||||
|
||||
|
||||
|
||||
@@ -8,8 +8,12 @@ diffusion models.
|
||||
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import threading
|
||||
import time
|
||||
import tempfile
|
||||
import warnings
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
@@ -18,9 +22,19 @@ import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from fastvideo.api.compat import (
|
||||
expand_request_prompt_batch,
|
||||
generator_config_to_fastvideo_args,
|
||||
legacy_from_pretrained_to_config,
|
||||
load_generator_config_from_file,
|
||||
normalize_generation_request,
|
||||
normalize_generator_config,
|
||||
request_to_pipeline_overrides,
|
||||
request_to_sampling_param,
|
||||
)
|
||||
from fastvideo.api.results import GenerationResult
|
||||
from fastvideo.api.schema import GenerationRequest, GeneratorConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -30,6 +44,29 @@ from fastvideo.worker.executor import Executor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_FROM_PRETRAINED_CONVENIENCE_KWARGS = frozenset({
|
||||
"num_gpus",
|
||||
"revision",
|
||||
"trust_remote_code",
|
||||
"distributed_executor_backend",
|
||||
"tp_size",
|
||||
"sp_size",
|
||||
"hsdp_replicate_dim",
|
||||
"hsdp_shard_dim",
|
||||
"dist_timeout",
|
||||
"use_fsdp_inference",
|
||||
"disable_autocast",
|
||||
"enable_stage_verification",
|
||||
"dit_cpu_offload",
|
||||
"dit_layerwise_offload",
|
||||
"text_encoder_cpu_offload",
|
||||
"image_encoder_cpu_offload",
|
||||
"vae_cpu_offload",
|
||||
"pin_cpu_memory",
|
||||
"enable_torch_compile",
|
||||
"torch_compile_kwargs",
|
||||
})
|
||||
|
||||
|
||||
def _infer_latent_batch_size(batch: ForwardBatch) -> int:
|
||||
if isinstance(batch.prompt, list):
|
||||
@@ -52,19 +89,33 @@ class VideoGenerator:
|
||||
customization options, similar to popular frameworks like HF Diffusers.
|
||||
"""
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, executor_class: type[Executor], log_stats: bool):
|
||||
def __init__(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
executor_class: type[Executor],
|
||||
log_stats: bool,
|
||||
*,
|
||||
log_queue=None,
|
||||
):
|
||||
"""
|
||||
Initialize the video generator.
|
||||
|
||||
|
||||
Args:
|
||||
fastvideo_args: The inference arguments
|
||||
executor_class: The executor class to use for inference
|
||||
log_stats: Whether to log statistics
|
||||
log_queue: Optional multiprocessing.Queue to forward worker logs to
|
||||
"""
|
||||
self.config: GeneratorConfig | None = None
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self.executor = executor_class(fastvideo_args)
|
||||
self.executor = executor_class(fastvideo_args, log_queue=log_queue)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str, **kwargs) -> "VideoGenerator":
|
||||
def from_pretrained(
|
||||
cls,
|
||||
model_path: str | GeneratorConfig | Mapping[str, Any] | None = None,
|
||||
**kwargs,
|
||||
) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
|
||||
@@ -77,21 +128,84 @@ class VideoGenerator:
|
||||
The created video generator
|
||||
|
||||
Priority level: Default pipeline config < User's pipeline config < User's kwargs
|
||||
"""
|
||||
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
|
||||
|
||||
return cls.from_fastvideo_args(fastvideo_args)
|
||||
Stable convenience kwargs remain supported here for common engine and
|
||||
offload settings. Advanced model- or pipeline-specific options should
|
||||
move to VideoGenerator.from_config(...).
|
||||
"""
|
||||
log_queue = kwargs.pop("log_queue", None)
|
||||
typed_config = kwargs.pop("config", None)
|
||||
if typed_config is not None:
|
||||
if model_path is not None:
|
||||
raise TypeError("Pass either model_path or config to from_pretrained, not both")
|
||||
if kwargs:
|
||||
unexpected = ", ".join(sorted(kwargs))
|
||||
raise TypeError(f"Unexpected keyword arguments with config: {unexpected}")
|
||||
return cls.from_config(typed_config, log_queue=log_queue)
|
||||
|
||||
if isinstance(model_path, GeneratorConfig | Mapping):
|
||||
if kwargs:
|
||||
unexpected = ", ".join(sorted(kwargs))
|
||||
raise TypeError(f"Unexpected keyword arguments with typed config: {unexpected}")
|
||||
return cls.from_config(model_path, log_queue=log_queue)
|
||||
|
||||
if model_path is None:
|
||||
raise TypeError("model_path or config is required")
|
||||
|
||||
legacy_only_kwargs = sorted(set(kwargs) - _FROM_PRETRAINED_CONVENIENCE_KWARGS)
|
||||
if legacy_only_kwargs:
|
||||
warnings.warn(
|
||||
"VideoGenerator.from_pretrained(...) received legacy-only kwargs "
|
||||
f"({', '.join(legacy_only_kwargs)}); prefer VideoGenerator.from_config(...) "
|
||||
"for advanced configuration.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
return cls.from_config(
|
||||
legacy_from_pretrained_to_config(model_path, kwargs),
|
||||
log_queue=log_queue,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_fastvideo_args(cls, fastvideo_args: FastVideoArgs) -> "VideoGenerator":
|
||||
def from_config(
|
||||
cls,
|
||||
config: GeneratorConfig | Mapping[str, Any],
|
||||
*,
|
||||
log_queue=None,
|
||||
) -> "VideoGenerator":
|
||||
normalized = normalize_generator_config(config)
|
||||
fastvideo_args = generator_config_to_fastvideo_args(normalized)
|
||||
generator = cls.from_fastvideo_args(fastvideo_args, log_queue=log_queue)
|
||||
generator.config = normalized
|
||||
return generator
|
||||
|
||||
@classmethod
|
||||
def from_file(
|
||||
cls,
|
||||
path: str,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
*,
|
||||
log_queue=None,
|
||||
) -> "VideoGenerator":
|
||||
return cls.from_config(
|
||||
load_generator_config_from_file(path, overrides=overrides),
|
||||
log_queue=log_queue,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_fastvideo_args(
|
||||
cls,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
*,
|
||||
log_queue=None,
|
||||
) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator with the specified arguments.
|
||||
|
||||
|
||||
Args:
|
||||
fastvideo_args: The inference arguments
|
||||
|
||||
log_queue: Optional multiprocessing.Queue to forward worker logs to
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
"""
|
||||
@@ -103,8 +217,40 @@ class VideoGenerator:
|
||||
fastvideo_args=fastvideo_args,
|
||||
executor_class=executor_class,
|
||||
log_stats=False, # TODO: implement
|
||||
log_queue=log_queue,
|
||||
)
|
||||
|
||||
def generate(
|
||||
self,
|
||||
request: GenerationRequest | Mapping[str, Any],
|
||||
*,
|
||||
log_queue=None,
|
||||
) -> GenerationResult | list[GenerationResult]:
|
||||
"""
|
||||
Generate video or image outputs from a typed inference request.
|
||||
|
||||
Args:
|
||||
request: A `GenerationRequest` instance or a mapping that can be
|
||||
parsed into one. This is the primary public inference
|
||||
entrypoint for the typed API.
|
||||
log_queue: Optional multiprocessing.Queue to forward worker logs to
|
||||
during this request.
|
||||
|
||||
Returns:
|
||||
A `GenerationResult` for single-request generation, or a list of
|
||||
`GenerationResult` objects when the request expands into multiple
|
||||
prompts.
|
||||
"""
|
||||
normalized_request = normalize_generation_request(request)
|
||||
if log_queue:
|
||||
self.executor.set_log_queue(log_queue)
|
||||
|
||||
try:
|
||||
return self._generate_request_impl(normalized_request)
|
||||
finally:
|
||||
if log_queue:
|
||||
self.executor.clear_log_queue()
|
||||
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str | None = None,
|
||||
@@ -140,9 +286,93 @@ class VideoGenerator:
|
||||
A metadata dictionary for single-prompt generation, or a list of
|
||||
metadata dictionaries for prompt-file batch generation.
|
||||
"""
|
||||
log_queue = kwargs.pop("log_queue", None)
|
||||
warnings.warn(
|
||||
"VideoGenerator.generate_video(...) is deprecated; use "
|
||||
"VideoGenerator.generate(request=...) instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if log_queue:
|
||||
self.executor.set_log_queue(log_queue)
|
||||
|
||||
try:
|
||||
return self._generate_video_impl(
|
||||
prompt=prompt,
|
||||
sampling_param=sampling_param,
|
||||
mouse_cond=mouse_cond,
|
||||
keyboard_cond=keyboard_cond,
|
||||
grid_sizes=grid_sizes,
|
||||
**kwargs,
|
||||
)
|
||||
finally:
|
||||
if log_queue:
|
||||
self.executor.clear_log_queue()
|
||||
|
||||
def _generate_request_impl(
|
||||
self,
|
||||
request: GenerationRequest,
|
||||
) -> GenerationResult | list[GenerationResult]:
|
||||
if isinstance(request.prompt, list):
|
||||
if request.inputs.prompt_path is not None:
|
||||
raise ValueError("request.prompt list cannot be combined with request.inputs.prompt_path")
|
||||
results: list[GenerationResult] = []
|
||||
for index, single_request in enumerate(expand_request_prompt_batch(request)):
|
||||
prompt = single_request.prompt
|
||||
wrapped = self._generate_single_request(single_request)
|
||||
if isinstance(wrapped, list):
|
||||
results.extend(wrapped)
|
||||
continue
|
||||
wrapped.prompt_index = index
|
||||
if wrapped.prompt is None:
|
||||
wrapped.prompt = prompt
|
||||
results.append(wrapped)
|
||||
return results
|
||||
|
||||
return self._generate_single_request(request)
|
||||
|
||||
def _generate_single_request(
|
||||
self,
|
||||
request: GenerationRequest,
|
||||
) -> GenerationResult | list[GenerationResult]:
|
||||
fastvideo_args = self.fastvideo_args
|
||||
pipeline_overrides = request_to_pipeline_overrides(request)
|
||||
if pipeline_overrides:
|
||||
fastvideo_args = deepcopy(self.fastvideo_args)
|
||||
for key, value in pipeline_overrides.items():
|
||||
if not hasattr(fastvideo_args.pipeline_config, key):
|
||||
raise ValueError(f"Request field {key!r} is not supported by pipeline config overrides")
|
||||
setattr(fastvideo_args.pipeline_config, key, deepcopy(value))
|
||||
|
||||
sampling_param = request_to_sampling_param(
|
||||
request,
|
||||
model_path=self.fastvideo_args.model_path,
|
||||
)
|
||||
result = self._generate_video_impl(
|
||||
prompt=request.prompt,
|
||||
sampling_param=sampling_param,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
return self._wrap_legacy_result(result)
|
||||
|
||||
def _generate_video_impl(
|
||||
self,
|
||||
prompt: str | None = None,
|
||||
sampling_param: SamplingParam | None = None,
|
||||
mouse_cond: torch.Tensor | None = None,
|
||||
keyboard_cond: torch.Tensor | None = None,
|
||||
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
|
||||
| None = None,
|
||||
fastvideo_args: FastVideoArgs | None = None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
|
||||
"""Internal implementation of generate_video."""
|
||||
if fastvideo_args is None:
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
# Handle batch processing from text file
|
||||
if sampling_param is None:
|
||||
sampling_param = SamplingParam.from_pretrained(self.fastvideo_args.model_path)
|
||||
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
|
||||
|
||||
# Add action control inputs to kwargs if provided
|
||||
if mouse_cond is not None:
|
||||
@@ -154,9 +384,9 @@ class VideoGenerator:
|
||||
|
||||
sampling_param.update(kwargs)
|
||||
|
||||
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
|
||||
prompt_txt_path = sampling_param.prompt_path or self.fastvideo_args.prompt_txt
|
||||
if not os.path.exists(prompt_txt_path):
|
||||
if fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
|
||||
prompt_txt_path = sampling_param.prompt_path or fastvideo_args.prompt_txt
|
||||
if not prompt_txt_path or not os.path.exists(prompt_txt_path):
|
||||
raise FileNotFoundError(f"Prompt text file not found: {prompt_txt_path}")
|
||||
|
||||
# Read prompts from file
|
||||
@@ -175,7 +405,12 @@ class VideoGenerator:
|
||||
# Generate video for this prompt using the same logic below
|
||||
output_path = self._prepare_output_path(sampling_param.output_path, batch_prompt)
|
||||
kwargs["output_path"] = output_path
|
||||
result = self._generate_single_video(prompt=batch_prompt, sampling_param=sampling_param, **kwargs)
|
||||
result = self._generate_single_video(
|
||||
prompt=batch_prompt,
|
||||
sampling_param=sampling_param,
|
||||
fastvideo_args=fastvideo_args,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Add prompt info to result
|
||||
result["prompt_index"] = i
|
||||
@@ -196,7 +431,12 @@ 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
|
||||
return self._generate_single_video(prompt=prompt, sampling_param=sampling_param, **kwargs)
|
||||
return self._generate_single_video(
|
||||
prompt=prompt,
|
||||
sampling_param=sampling_param,
|
||||
fastvideo_args=fastvideo_args,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def _is_image_workload(self) -> bool:
|
||||
"""Return True when the workload produces a single image (t2i, i2i …)."""
|
||||
@@ -280,11 +520,12 @@ class VideoGenerator:
|
||||
self,
|
||||
prompt: str,
|
||||
sampling_param: SamplingParam | None = None,
|
||||
fastvideo_args: FastVideoArgs | None = None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any]:
|
||||
"""Internal method for single video generation"""
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args
|
||||
if fastvideo_args is None:
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
@@ -427,6 +668,20 @@ class VideoGenerator:
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _wrap_legacy_result(
|
||||
result: dict[str, Any] | list[dict[str, Any]], ) -> GenerationResult | list[GenerationResult]:
|
||||
if isinstance(result, list):
|
||||
return [GenerationResult.from_legacy_result(item) for item in result]
|
||||
return GenerationResult.from_legacy_result(result)
|
||||
|
||||
@staticmethod
|
||||
def _unwrap_typed_result(
|
||||
result: GenerationResult | list[GenerationResult], ) -> dict[str, Any] | list[dict[str, Any]]:
|
||||
if isinstance(result, list):
|
||||
return [item.to_legacy_dict() for item in result]
|
||||
return result.to_legacy_dict()
|
||||
|
||||
@staticmethod
|
||||
def _mux_audio(
|
||||
video_path: str,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -588,6 +588,33 @@ class TokenizerLoader(ComponentLoader):
|
||||
class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
|
||||
@staticmethod
|
||||
def _find_gen3c_tokenizer_checkpoint(model_path: str) -> str | None:
|
||||
"""Locate tokenizer-backed VAE checkpoint used by GEN3C integration."""
|
||||
candidates = [
|
||||
os.path.join(model_path, "tokenizer.pth"),
|
||||
os.path.join(os.path.dirname(model_path), "tokenizer",
|
||||
"tokenizer.pth"),
|
||||
]
|
||||
for candidate in candidates:
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _find_gen3c_jit_tokenizer_dir(model_path: str) -> str | None:
|
||||
"""Locate official tokenizer JIT assets (encoder/decoder/mean_std)."""
|
||||
candidates = [
|
||||
model_path,
|
||||
os.path.join(os.path.dirname(model_path), "tokenizer"),
|
||||
]
|
||||
required = ("encoder.jit", "decoder.jit", "mean_std.pt")
|
||||
for directory in candidates:
|
||||
if all(os.path.exists(os.path.join(directory, name))
|
||||
for name in required):
|
||||
return directory
|
||||
return None
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
@@ -614,8 +641,68 @@ class VAELoader(ComponentLoader):
|
||||
if fastvideo_args.pipeline_config.vae_precision
|
||||
else torch.bfloat16
|
||||
):
|
||||
pipeline_name = fastvideo_args.pipeline_config.__class__.__name__
|
||||
is_gen3c = pipeline_name.startswith("Gen3C")
|
||||
is_cosmos25 = pipeline_name == "Cosmos25Config"
|
||||
|
||||
# GEN3C: prefer tokenizer-backed VAE checkpoint when available.
|
||||
# This aligns latent conditioning with the GEN3C temporal contract.
|
||||
if is_gen3c and class_name in (
|
||||
"AutoencoderKLWan", "AutoencoderKLGen3CTokenizer"):
|
||||
from fastvideo.models.vaes.gen3c_tokenizer_vae import (
|
||||
AutoencoderKLGen3CTokenizer)
|
||||
|
||||
dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
num_frames = int(
|
||||
getattr(fastvideo_args.pipeline_config, "num_frames", 121))
|
||||
state_t = int(
|
||||
getattr(fastvideo_args.pipeline_config, "state_t", 16))
|
||||
if state_t > 1 and num_frames > 1:
|
||||
target_temporal = max(1,
|
||||
(num_frames - 1) // (state_t - 1))
|
||||
else:
|
||||
target_temporal = 8
|
||||
|
||||
jit_dir = self._find_gen3c_jit_tokenizer_dir(model_path)
|
||||
if jit_dir is not None:
|
||||
vae = AutoencoderKLGen3CTokenizer.from_jit_tokenizer(
|
||||
jit_dir,
|
||||
device=target_device,
|
||||
dtype=dtype,
|
||||
target_temporal_compression=target_temporal,
|
||||
pixel_chunk_duration=num_frames,
|
||||
)
|
||||
logger.info(
|
||||
"Loaded GEN3C tokenizer VAE from JIT assets in %s (target temporal compression=%d)",
|
||||
jit_dir,
|
||||
target_temporal,
|
||||
)
|
||||
return vae.eval()
|
||||
|
||||
tokenizer_ckpt = self._find_gen3c_tokenizer_checkpoint(
|
||||
model_path)
|
||||
if tokenizer_ckpt is not None:
|
||||
vae = AutoencoderKLGen3CTokenizer.from_tokenizer_checkpoint(
|
||||
tokenizer_ckpt,
|
||||
device=target_device,
|
||||
dtype=dtype,
|
||||
target_temporal_compression=target_temporal,
|
||||
pixel_chunk_duration=num_frames,
|
||||
)
|
||||
logger.info(
|
||||
"Loaded GEN3C tokenizer VAE from %s (target temporal compression=%d)",
|
||||
tokenizer_ckpt,
|
||||
target_temporal,
|
||||
)
|
||||
return vae.eval()
|
||||
logger.warning(
|
||||
"GEN3C tokenizer VAE checkpoint not found near %s; falling back to configured class %s.",
|
||||
model_path,
|
||||
class_name,
|
||||
)
|
||||
|
||||
# Cosmos2.5 uses a Wan2.1 VAE stored as `tokenizer.safetensors` under the VAE folder.
|
||||
is_cosmos25 = fastvideo_args.pipeline_config.__class__.__name__ == "Cosmos25Config"
|
||||
if class_name == "AutoencoderKLWan" and is_cosmos25:
|
||||
from fastvideo.models.vaes.cosmos25wanvae import Cosmos25WanVAE
|
||||
|
||||
|
||||
@@ -40,8 +40,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
|
||||
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
|
||||
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
|
||||
"Kandinsky5Transformer3DModel":
|
||||
("dits", "kandinsky5", "Kandinsky5Transformer3DModel"),
|
||||
"Gen3CTransformer3DModel": ("dits", "gen3c", "Gen3CTransformer3DModel"),
|
||||
"Kandinsky5Transformer3DModel": ("dits", "kandinsky5", "Kandinsky5Transformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -82,6 +82,9 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
|
||||
"AutoencoderKLGen3CTokenizer":
|
||||
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
|
||||
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,366 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
GEN3C tokenizer-backed VAE adapter.
|
||||
|
||||
This wrapper loads the available tokenizer checkpoint (`tokenizer.pth`) and
|
||||
adapts it to GEN3C's latent-time contract (T=16 for 121 output frames).
|
||||
|
||||
Why this exists:
|
||||
- The converted GEN3C bundle includes tokenizer-style VAE weights, not a
|
||||
standard diffusers Wan VAE contract.
|
||||
- GEN3C diffusion expects 8x temporal compression (121 -> 16), while the
|
||||
available tokenizer checkpoint follows a 4x temporal path.
|
||||
|
||||
To bridge this at inference time, we:
|
||||
- keep the inner tokenizer model as-is,
|
||||
- downsample encoded latent time from inner-T to target-T for DiT input,
|
||||
- upsample generated latent time back to inner-T before decoding.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _TensorLatentDist:
|
||||
"""Minimal distribution-like wrapper used by pipeline stages."""
|
||||
|
||||
mean: torch.Tensor
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self.mean
|
||||
|
||||
def sample(self, generator: Any | None = None) -> torch.Tensor:
|
||||
_ = generator
|
||||
return self.mean
|
||||
|
||||
|
||||
class _JITGen3CTokenizerInner(nn.Module):
|
||||
"""Minimal wrapper around official tokenizer JIT encoder/decoder exports."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
encoder_path: str,
|
||||
decoder_path: str,
|
||||
mean_std_path: str,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
latent_channels: int = 16,
|
||||
latent_chunk_duration: int = 16,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._dtype = dtype
|
||||
self._forced_bf16 = False
|
||||
self.encoder = torch.jit.load(encoder_path, map_location=device).eval().to(
|
||||
device=device, dtype=dtype)
|
||||
self.decoder = torch.jit.load(decoder_path, map_location=device).eval().to(
|
||||
device=device, dtype=dtype)
|
||||
|
||||
latent_mean, latent_std = torch.load(mean_std_path, map_location="cpu")
|
||||
latent_mean = latent_mean.view(latent_channels, -1)[:, :latent_chunk_duration]
|
||||
latent_std = latent_std.view(latent_channels, -1)[:, :latent_chunk_duration]
|
||||
|
||||
self.register_buffer(
|
||||
"_latent_mean",
|
||||
latent_mean.to(torch.float32).view(1, latent_channels,
|
||||
latent_chunk_duration, 1, 1),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer(
|
||||
"_latent_std",
|
||||
latent_std.to(torch.float32).view(1, latent_channels,
|
||||
latent_chunk_duration, 1, 1),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
def _match_stats(self, like: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
mean = self._latent_mean.to(device=like.device, dtype=like.dtype)
|
||||
std = self._latent_std.to(device=like.device, dtype=like.dtype)
|
||||
t = like.shape[2]
|
||||
if mean.shape[2] == t:
|
||||
return mean, std
|
||||
if t < mean.shape[2]:
|
||||
return mean[:, :, :t], std[:, :, :t]
|
||||
# fallback for non-default lengths
|
||||
mean = torch.nn.functional.interpolate(
|
||||
mean, size=(t, 1, 1), mode="trilinear", align_corners=False)
|
||||
std = torch.nn.functional.interpolate(
|
||||
std, size=(t, 1, 1), mode="trilinear", align_corners=False)
|
||||
return mean, std
|
||||
|
||||
@staticmethod
|
||||
def _module_dtype_device(module: torch.nn.Module) -> tuple[torch.dtype, torch.device]:
|
||||
for param in module.parameters():
|
||||
return param.dtype, param.device
|
||||
for buf in module.buffers():
|
||||
return buf.dtype, buf.device
|
||||
raise RuntimeError("Tokenizer JIT module has no parameters/buffers to infer dtype/device.")
|
||||
|
||||
def _coerce_modules_to_bf16(self) -> None:
|
||||
if self._forced_bf16:
|
||||
return
|
||||
self.encoder = self.encoder.to(dtype=torch.bfloat16)
|
||||
self.decoder = self.decoder.to(dtype=torch.bfloat16)
|
||||
self._dtype = torch.bfloat16
|
||||
self._forced_bf16 = True
|
||||
logger.warning(
|
||||
"GEN3C tokenizer JIT hit fp16/bf16 mismatch; coercing tokenizer encoder/decoder to bf16."
|
||||
)
|
||||
|
||||
def encode(self, x: torch.Tensor) -> _TensorLatentDist:
|
||||
enc_dtype, enc_device = self._module_dtype_device(self.encoder)
|
||||
x_in = x.to(device=enc_device, dtype=enc_dtype)
|
||||
try:
|
||||
with torch.autocast(device_type=enc_device.type, enabled=False):
|
||||
z = self.encoder(x_in)
|
||||
except RuntimeError as e:
|
||||
err = str(e)
|
||||
mismatch_tokens = (
|
||||
"Input type (CUDABFloat16Type) and weight type (torch.cuda.HalfTensor)",
|
||||
"Input type (torch.cuda.HalfTensor) and weight type (CUDABFloat16Type)",
|
||||
)
|
||||
if any(token in err for token in mismatch_tokens):
|
||||
self._coerce_modules_to_bf16()
|
||||
enc_dtype, enc_device = self._module_dtype_device(self.encoder)
|
||||
x_in = x.to(device=enc_device, dtype=enc_dtype)
|
||||
with torch.autocast(device_type=enc_device.type, enabled=False):
|
||||
z = self.encoder(x_in)
|
||||
else:
|
||||
raise
|
||||
if isinstance(z, tuple):
|
||||
z = z[0]
|
||||
z = z.to(dtype=x.dtype, device=x.device)
|
||||
mean, std = self._match_stats(z)
|
||||
return _TensorLatentDist((z - mean) / std)
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
mean, std = self._match_stats(z)
|
||||
dec_dtype, dec_device = self._module_dtype_device(self.decoder)
|
||||
z_in = (z * std + mean).to(device=dec_device, dtype=dec_dtype)
|
||||
with torch.autocast(device_type=dec_device.type, enabled=False):
|
||||
x = self.decoder(z_in)
|
||||
if isinstance(x, tuple):
|
||||
x = x[0]
|
||||
return x.to(dtype=z.dtype, device=z.device)
|
||||
|
||||
|
||||
class AutoencoderKLGen3CTokenizer(nn.Module):
|
||||
"""
|
||||
GEN3C VAE wrapper with temporal contract adaptation.
|
||||
|
||||
Interface contract:
|
||||
- `encode(x)` returns normalized latents in the *target* temporal layout.
|
||||
- `decode(z)` expects normalized latents in the *target* temporal layout.
|
||||
"""
|
||||
|
||||
handles_latent_norm: bool = True
|
||||
handles_latent_denorm: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inner: nn.Module,
|
||||
*,
|
||||
target_temporal_compression: int = 8,
|
||||
inner_temporal_compression: int = 4,
|
||||
spatial_compression_factor: int = 8,
|
||||
pixel_chunk_duration: int = 121,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.inner = inner
|
||||
self.config = getattr(inner, "config", None)
|
||||
self._target_temporal_compression = int(target_temporal_compression)
|
||||
self._inner_temporal_compression = int(inner_temporal_compression)
|
||||
self._spatial_compression_factor = int(spatial_compression_factor)
|
||||
self._pixel_chunk_duration = int(pixel_chunk_duration)
|
||||
|
||||
@staticmethod
|
||||
def _extract_latents(encoder_output: Any) -> torch.Tensor:
|
||||
if hasattr(encoder_output, "latent_dist"):
|
||||
dist = encoder_output.latent_dist
|
||||
if hasattr(dist, "mode"):
|
||||
return dist.mode()
|
||||
if hasattr(dist, "mean"):
|
||||
return dist.mean
|
||||
return dist.sample()
|
||||
if hasattr(encoder_output, "mode"):
|
||||
return encoder_output.mode()
|
||||
if hasattr(encoder_output, "latents"):
|
||||
return encoder_output.latents
|
||||
if hasattr(encoder_output, "sample"):
|
||||
return encoder_output.sample()
|
||||
if isinstance(encoder_output, torch.Tensor):
|
||||
return encoder_output
|
||||
raise TypeError(f"Unsupported encoder output type: {type(encoder_output)}")
|
||||
|
||||
def _inner_to_target_time(self, z_inner: torch.Tensor) -> torch.Tensor:
|
||||
if z_inner.shape[2] <= 1:
|
||||
return z_inner
|
||||
|
||||
# Common GEN3C case: inner=4x, target=8x => keep every other latent frame.
|
||||
if self._target_temporal_compression == 2 * self._inner_temporal_compression:
|
||||
return z_inner[:, :, 0::2, :, :].contiguous()
|
||||
|
||||
# Generic fallback: keep boundary latents and sample uniformly.
|
||||
t_inner = z_inner.shape[2]
|
||||
t_target = 1 + (t_inner - 1) * self._inner_temporal_compression // self._target_temporal_compression
|
||||
idx = torch.linspace(0, t_inner - 1, t_target, device=z_inner.device)
|
||||
idx = idx.round().long()
|
||||
return z_inner.index_select(2, idx).contiguous()
|
||||
|
||||
def _target_to_inner_time(self, z_target: torch.Tensor) -> torch.Tensor:
|
||||
if z_target.shape[2] <= 1:
|
||||
return z_target
|
||||
|
||||
# Common GEN3C case: inner=4x, target=8x => insert midpoint frames.
|
||||
if self._target_temporal_compression == 2 * self._inner_temporal_compression:
|
||||
b, c, t, h, w = z_target.shape
|
||||
t_inner = 2 * t - 1
|
||||
out = torch.empty(
|
||||
b, c, t_inner, h, w, device=z_target.device, dtype=z_target.dtype)
|
||||
out[:, :, 0::2, :, :] = z_target
|
||||
out[:, :, 1::2, :, :] = 0.5 * (
|
||||
z_target[:, :, :-1, :, :] + z_target[:, :, 1:, :, :]
|
||||
)
|
||||
return out.contiguous()
|
||||
|
||||
# Generic fallback: linear index interpolation in time.
|
||||
t_target = z_target.shape[2]
|
||||
t_inner = 1 + (t_target - 1) * self._target_temporal_compression // self._inner_temporal_compression
|
||||
idx = torch.linspace(0, t_target - 1, t_inner, device=z_target.device)
|
||||
idx0 = idx.floor().long()
|
||||
idx1 = idx.ceil().long().clamp_max(t_target - 1)
|
||||
frac = (idx - idx0).view(1, 1, -1, 1, 1)
|
||||
z0 = z_target.index_select(2, idx0)
|
||||
z1 = z_target.index_select(2, idx1)
|
||||
return (z0 * (1.0 - frac) + z1 * frac).contiguous()
|
||||
|
||||
def encode(self, x: torch.Tensor) -> _TensorLatentDist:
|
||||
z_inner = self._extract_latents(self.inner.encode(x))
|
||||
z_target = self._inner_to_target_time(z_inner)
|
||||
return _TensorLatentDist(z_target)
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
z_inner = self._target_to_inner_time(z)
|
||||
out = self.inner.decode(z_inner)
|
||||
return out.sample if hasattr(out, "sample") else out
|
||||
|
||||
def enable_tiling(self) -> None:
|
||||
if hasattr(self.inner, "enable_tiling"):
|
||||
self.inner.enable_tiling()
|
||||
|
||||
def disable_tiling(self) -> None:
|
||||
if hasattr(self.inner, "disable_tiling"):
|
||||
self.inner.disable_tiling()
|
||||
|
||||
def get_latent_num_frames(self, num_pixel_frames: int) -> int:
|
||||
num_pixel_frames = int(num_pixel_frames)
|
||||
if num_pixel_frames <= 1:
|
||||
return 1
|
||||
return 1 + (num_pixel_frames - 1) // self._target_temporal_compression
|
||||
|
||||
def get_pixel_num_frames(self, num_latent_frames: int) -> int:
|
||||
num_latent_frames = int(num_latent_frames)
|
||||
if num_latent_frames <= 1:
|
||||
return 1
|
||||
return (num_latent_frames - 1) * self._target_temporal_compression + 1
|
||||
|
||||
@property
|
||||
def spatial_compression_factor(self) -> int:
|
||||
return self._spatial_compression_factor
|
||||
|
||||
@property
|
||||
def temporal_compression_factor(self) -> int:
|
||||
return self._target_temporal_compression
|
||||
|
||||
@property
|
||||
def temporal_compression_ratio(self) -> int:
|
||||
return self._target_temporal_compression
|
||||
|
||||
@property
|
||||
def pixel_chunk_duration(self) -> int:
|
||||
return self._pixel_chunk_duration
|
||||
|
||||
@property
|
||||
def latent_chunk_duration(self) -> int:
|
||||
return self.get_latent_num_frames(self._pixel_chunk_duration)
|
||||
|
||||
@classmethod
|
||||
def from_tokenizer_checkpoint(
|
||||
cls,
|
||||
checkpoint_path: str,
|
||||
*,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
target_temporal_compression: int = 8,
|
||||
pixel_chunk_duration: int = 121,
|
||||
) -> "AutoencoderKLGen3CTokenizer":
|
||||
from fastvideo.models.vaes.cosmos25wanvae import Cosmos25WanVAE
|
||||
|
||||
inner = Cosmos25WanVAE(device=device, dtype=dtype)
|
||||
loaded = torch.load(checkpoint_path, map_location="cpu")
|
||||
if isinstance(loaded, dict):
|
||||
for key in ("state_dict", "model", "ema", "model_state_dict"):
|
||||
if key in loaded and isinstance(loaded[key], dict):
|
||||
loaded = loaded[key]
|
||||
break
|
||||
missing, unexpected = inner.load_state_dict(loaded, strict=False)
|
||||
if missing:
|
||||
logger.warning(
|
||||
"GEN3C tokenizer VAE missing keys (%d). Example: %s",
|
||||
len(missing),
|
||||
missing[:5],
|
||||
)
|
||||
if unexpected:
|
||||
logger.warning(
|
||||
"GEN3C tokenizer VAE unexpected keys (%d). Example: %s",
|
||||
len(unexpected),
|
||||
unexpected[:5],
|
||||
)
|
||||
return cls(
|
||||
inner,
|
||||
target_temporal_compression=target_temporal_compression,
|
||||
inner_temporal_compression=4,
|
||||
spatial_compression_factor=8,
|
||||
pixel_chunk_duration=pixel_chunk_duration,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_jit_tokenizer(
|
||||
cls,
|
||||
tokenizer_dir: str,
|
||||
*,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
target_temporal_compression: int = 8,
|
||||
pixel_chunk_duration: int = 121,
|
||||
) -> "AutoencoderKLGen3CTokenizer":
|
||||
encoder_path = f"{tokenizer_dir}/encoder.jit"
|
||||
decoder_path = f"{tokenizer_dir}/decoder.jit"
|
||||
mean_std_path = f"{tokenizer_dir}/mean_std.pt"
|
||||
inner = _JITGen3CTokenizerInner(
|
||||
encoder_path=encoder_path,
|
||||
decoder_path=decoder_path,
|
||||
mean_std_path=mean_std_path,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
latent_channels=16,
|
||||
latent_chunk_duration=1 + (pixel_chunk_duration - 1) //
|
||||
target_temporal_compression,
|
||||
)
|
||||
return cls(
|
||||
inner,
|
||||
target_temporal_compression=target_temporal_compression,
|
||||
inner_temporal_compression=target_temporal_compression,
|
||||
spatial_compression_factor=8,
|
||||
pixel_chunk_duration=pixel_chunk_duration,
|
||||
)
|
||||
@@ -0,0 +1,41 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Sampling profiles for Cosmos and Cosmos 2.5 models."""
|
||||
from fastvideo.configs.sample.profiles import ModelProfile
|
||||
|
||||
_COSMOS_NEGATIVE_PROMPT = ("The video captures a series of frames showing ugly scenes, "
|
||||
"static with no motion, motion blur, over-saturation, "
|
||||
"shaky footage, low resolution, grainy texture, "
|
||||
"pixelated images, poorly lit areas, underexposed and "
|
||||
"overexposed scenes, poor color balance, washed out colors, "
|
||||
"choppy sequences, jerky movements, low frame rate, "
|
||||
"artifacting, color banding, unnatural transitions, "
|
||||
"outdated special effects, fake elements, unconvincing "
|
||||
"visuals, poorly edited content, jump cuts, visual noise, "
|
||||
"and flickering. Overall, the video is of poor quality.")
|
||||
|
||||
COSMOS_PREDICT2_2B = ModelProfile(
|
||||
name="cosmos_predict2_2b",
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 93,
|
||||
"fps": 16,
|
||||
"guidance_scale": 7.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": _COSMOS_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
COSMOS25_PREDICT2_2B = ModelProfile(
|
||||
name="cosmos25_predict2_2b",
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 77,
|
||||
"fps": 24,
|
||||
"seed": 0,
|
||||
"guidance_scale": 7.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": _COSMOS_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,36 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
GEN3C is a 3D-informed world-consistent video generation model with precise camera control.
|
||||
"""
|
||||
|
||||
from fastvideo.pipelines.basic.gen3c.cache_3d import (
|
||||
Cache3DBase,
|
||||
Cache3DBuffer,
|
||||
forward_warp,
|
||||
unproject_points,
|
||||
project_points,
|
||||
)
|
||||
from fastvideo.pipelines.basic.gen3c.gen3c_pipeline import (
|
||||
Gen3CPipeline,
|
||||
Gen3CConditioningStage,
|
||||
Gen3CDenoisingStage,
|
||||
Gen3CLatentPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.basic.gen3c.camera_utils import (
|
||||
generate_camera_trajectory, )
|
||||
|
||||
__all__ = [
|
||||
# 3D Cache
|
||||
"Cache3DBase",
|
||||
"Cache3DBuffer",
|
||||
"forward_warp",
|
||||
"unproject_points",
|
||||
"project_points",
|
||||
# Camera
|
||||
"generate_camera_trajectory",
|
||||
# Pipeline
|
||||
"Gen3CPipeline",
|
||||
"Gen3CConditioningStage",
|
||||
"Gen3CDenoisingStage",
|
||||
"Gen3CLatentPreparationStage",
|
||||
]
|
||||
@@ -0,0 +1,720 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
This module implements the 3D cache system for GEN3C video generation with camera control.
|
||||
The cache maintains a point cloud representation of the scene, enabling:
|
||||
- Unprojecting depth maps to 3D world points
|
||||
- Forward warping rendered views to new camera poses
|
||||
- Managing multiple frame buffers for temporal consistency
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def inverse_with_conversion(mtx: torch.Tensor) -> torch.Tensor:
|
||||
"""Compute matrix inverse with float32 conversion for numerical stability."""
|
||||
return torch.linalg.inv(mtx.to(torch.float32)).to(mtx.dtype)
|
||||
|
||||
|
||||
def create_grid(b: int, h: int, w: int, device: str = "cpu", dtype: torch.dtype = torch.float32) -> torch.Tensor:
|
||||
"""
|
||||
Create a dense grid of (x, y) coordinates of shape (b, 2, h, w).
|
||||
|
||||
Args:
|
||||
b: Batch size
|
||||
h: Height
|
||||
w: Width
|
||||
device: Device for tensor creation
|
||||
dtype: Data type for tensor
|
||||
|
||||
Returns:
|
||||
Grid tensor of shape (b, 2, h, w)
|
||||
"""
|
||||
x = torch.arange(0, w, device=device, dtype=dtype).view(1, 1, 1, w).expand(b, 1, h, w)
|
||||
y = torch.arange(0, h, device=device, dtype=dtype).view(1, 1, h, 1).expand(b, 1, h, w)
|
||||
return torch.cat([x, y], dim=1)
|
||||
|
||||
|
||||
def unproject_points(
|
||||
depth: torch.Tensor,
|
||||
w2c: torch.Tensor,
|
||||
intrinsic: torch.Tensor,
|
||||
is_depth: bool = True,
|
||||
mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Unproject depth map to 3D world points.
|
||||
|
||||
Args:
|
||||
depth: (b, 1, h, w) depth map
|
||||
w2c: (b, 4, 4) world-to-camera transformation matrix
|
||||
intrinsic: (b, 3, 3) camera intrinsic matrix
|
||||
is_depth: If True, depth is z-depth; if False, depth is distance to camera
|
||||
mask: Optional (b, h, w) or (b, 1, h, w) mask for valid pixels
|
||||
|
||||
Returns:
|
||||
world_points: (b, h, w, 3) 3D world coordinates
|
||||
"""
|
||||
b, _, h, w = depth.shape
|
||||
device = depth.device
|
||||
dtype = depth.dtype
|
||||
|
||||
if mask is None:
|
||||
mask = depth > 0
|
||||
if mask.dim() == depth.dim() and mask.shape[1] == 1:
|
||||
mask = mask[:, 0]
|
||||
|
||||
idx = torch.nonzero(mask)
|
||||
if idx.numel() == 0:
|
||||
return torch.zeros((b, h, w, 3), device=device, dtype=dtype)
|
||||
|
||||
b_idx, y_idx, x_idx = idx[:, 0], idx[:, 1], idx[:, 2]
|
||||
|
||||
intrinsic_inv = inverse_with_conversion(intrinsic) # (b, 3, 3)
|
||||
|
||||
x_valid = x_idx.to(dtype)
|
||||
y_valid = y_idx.to(dtype)
|
||||
ones = torch.ones_like(x_valid)
|
||||
pos = torch.stack([x_valid, y_valid, ones], dim=1).unsqueeze(-1) # (N, 3, 1)
|
||||
|
||||
intrinsic_inv_valid = intrinsic_inv[b_idx] # (N, 3, 3)
|
||||
unnormalized_pos = torch.matmul(intrinsic_inv_valid, pos) # (N, 3, 1)
|
||||
|
||||
depth_valid = depth[b_idx, 0, y_idx, x_idx].view(-1, 1, 1)
|
||||
if is_depth:
|
||||
world_points_cam = depth_valid * unnormalized_pos
|
||||
else:
|
||||
norm_val = torch.norm(unnormalized_pos, dim=1, keepdim=True)
|
||||
direction = unnormalized_pos / (norm_val + 1e-8)
|
||||
world_points_cam = depth_valid * direction
|
||||
|
||||
ones_h = torch.ones((world_points_cam.shape[0], 1, 1), device=device, dtype=dtype)
|
||||
world_points_homo = torch.cat([world_points_cam, ones_h], dim=1) # (N, 4, 1)
|
||||
|
||||
trans = inverse_with_conversion(w2c) # (b, 4, 4)
|
||||
trans_valid = trans[b_idx] # (N, 4, 4)
|
||||
world_points_transformed = torch.matmul(trans_valid, world_points_homo) # (N, 4, 1)
|
||||
sparse_points = world_points_transformed[:, :3, 0] # (N, 3)
|
||||
|
||||
out_points = torch.zeros((b, h, w, 3), device=device, dtype=dtype)
|
||||
out_points[b_idx, y_idx, x_idx, :] = sparse_points
|
||||
return out_points
|
||||
|
||||
|
||||
def project_points(
|
||||
world_points: torch.Tensor,
|
||||
w2c: torch.Tensor,
|
||||
intrinsic: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Project 3D world points to 2D pixel coordinates.
|
||||
|
||||
Args:
|
||||
world_points: (b, h, w, 3) 3D world coordinates
|
||||
w2c: (b, 4, 4) world-to-camera transformation matrix
|
||||
intrinsic: (b, 3, 3) camera intrinsic matrix
|
||||
|
||||
Returns:
|
||||
projected_points: (b, h, w, 3, 1) projected 2D coordinates (x, y, z)
|
||||
"""
|
||||
world_points = world_points.unsqueeze(-1) # (b, h, w, 3, 1)
|
||||
b, h, w, _, _ = world_points.shape
|
||||
|
||||
ones_4d = torch.ones((b, h, w, 1, 1), device=world_points.device, dtype=world_points.dtype)
|
||||
world_points_homo = torch.cat([world_points, ones_4d], dim=3) # (b, h, w, 4, 1)
|
||||
|
||||
trans_4d = w2c[:, None, None] # (b, 1, 1, 4, 4)
|
||||
camera_points_homo = torch.matmul(trans_4d, world_points_homo) # (b, h, w, 4, 1)
|
||||
|
||||
camera_points = camera_points_homo[:, :, :, :3] # (b, h, w, 3, 1)
|
||||
intrinsic_4d = intrinsic[:, None, None] # (b, 1, 1, 3, 3)
|
||||
projected_points = torch.matmul(intrinsic_4d, camera_points) # (b, h, w, 3, 1)
|
||||
|
||||
return projected_points
|
||||
|
||||
|
||||
def bilinear_splatting(
|
||||
frame1: torch.Tensor,
|
||||
mask1: torch.Tensor | None,
|
||||
depth1: torch.Tensor,
|
||||
flow12: torch.Tensor,
|
||||
flow12_mask: torch.Tensor | None = None,
|
||||
is_image: bool = False,
|
||||
depth_weight_scale: float = 50.0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Bilinear splatting for forward warping.
|
||||
|
||||
Args:
|
||||
frame1: (b, c, h, w) source frame
|
||||
mask1: (b, 1, h, w) valid pixel mask (1 for known, 0 for unknown)
|
||||
depth1: (b, 1, h, w) depth map
|
||||
flow12: (b, 2, h, w) optical flow from frame1 to frame2
|
||||
flow12_mask: (b, 1, h, w) flow validity mask
|
||||
is_image: If True, output will be clipped to (-1, 1) range
|
||||
depth_weight_scale: Scale factor for depth weighting
|
||||
|
||||
Returns:
|
||||
warped_frame2: (b, c, h, w) warped frame
|
||||
mask2: (b, 1, h, w) validity mask for warped frame
|
||||
"""
|
||||
b, c, h, w = frame1.shape
|
||||
device = frame1.device
|
||||
dtype = frame1.dtype
|
||||
|
||||
if mask1 is None:
|
||||
mask1 = torch.ones(size=(b, 1, h, w), device=device, dtype=dtype)
|
||||
if flow12_mask is None:
|
||||
flow12_mask = torch.ones(size=(b, 1, h, w), device=device, dtype=dtype)
|
||||
|
||||
grid = create_grid(b, h, w, device=device, dtype=dtype)
|
||||
trans_pos = flow12 + grid
|
||||
|
||||
trans_pos_offset = trans_pos + 1
|
||||
trans_pos_floor = torch.floor(trans_pos_offset).long()
|
||||
trans_pos_ceil = torch.ceil(trans_pos_offset).long()
|
||||
|
||||
trans_pos_offset = torch.stack(
|
||||
[torch.clamp(trans_pos_offset[:, 0], min=0, max=w + 1),
|
||||
torch.clamp(trans_pos_offset[:, 1], min=0, max=h + 1)],
|
||||
dim=1)
|
||||
trans_pos_floor = torch.stack(
|
||||
[torch.clamp(trans_pos_floor[:, 0], min=0, max=w + 1),
|
||||
torch.clamp(trans_pos_floor[:, 1], min=0, max=h + 1)],
|
||||
dim=1)
|
||||
trans_pos_ceil = torch.stack(
|
||||
[torch.clamp(trans_pos_ceil[:, 0], min=0, max=w + 1),
|
||||
torch.clamp(trans_pos_ceil[:, 1], min=0, max=h + 1)],
|
||||
dim=1)
|
||||
|
||||
# Bilinear weights
|
||||
prox_weight_nw = (1 - (trans_pos_offset[:, 1:2] - trans_pos_floor[:, 1:2])) * \
|
||||
(1 - (trans_pos_offset[:, 0:1] - trans_pos_floor[:, 0:1]))
|
||||
prox_weight_sw = (1 - (trans_pos_ceil[:, 1:2] - trans_pos_offset[:, 1:2])) * \
|
||||
(1 - (trans_pos_offset[:, 0:1] - trans_pos_floor[:, 0:1]))
|
||||
prox_weight_ne = (1 - (trans_pos_offset[:, 1:2] - trans_pos_floor[:, 1:2])) * \
|
||||
(1 - (trans_pos_ceil[:, 0:1] - trans_pos_offset[:, 0:1]))
|
||||
prox_weight_se = (1 - (trans_pos_ceil[:, 1:2] - trans_pos_offset[:, 1:2])) * \
|
||||
(1 - (trans_pos_ceil[:, 0:1] - trans_pos_offset[:, 0:1]))
|
||||
|
||||
# Depth weighting for occlusion handling
|
||||
clamped_depth1 = torch.clamp(depth1, min=0)
|
||||
log_depth1 = torch.log1p(clamped_depth1)
|
||||
exponent = log_depth1 / (log_depth1.max() + 1e-7) * depth_weight_scale
|
||||
max_exponent = 80.0 if dtype in [torch.float32, torch.bfloat16] else 10.0
|
||||
clamped_exponent = torch.clamp(exponent, max=max_exponent)
|
||||
depth_weights = torch.exp(clamped_exponent) + 1e-7
|
||||
|
||||
weight_nw = torch.moveaxis(prox_weight_nw * mask1 * flow12_mask / depth_weights, [0, 1, 2, 3], [0, 3, 1, 2])
|
||||
weight_sw = torch.moveaxis(prox_weight_sw * mask1 * flow12_mask / depth_weights, [0, 1, 2, 3], [0, 3, 1, 2])
|
||||
weight_ne = torch.moveaxis(prox_weight_ne * mask1 * flow12_mask / depth_weights, [0, 1, 2, 3], [0, 3, 1, 2])
|
||||
weight_se = torch.moveaxis(prox_weight_se * mask1 * flow12_mask / depth_weights, [0, 1, 2, 3], [0, 3, 1, 2])
|
||||
|
||||
warped_frame = torch.zeros(size=(b, h + 2, w + 2, c), dtype=dtype, device=device)
|
||||
warped_weights = torch.zeros(size=(b, h + 2, w + 2, 1), dtype=dtype, device=device)
|
||||
|
||||
frame1_cl = torch.moveaxis(frame1, [0, 1, 2, 3], [0, 3, 1, 2])
|
||||
batch_indices = torch.arange(b, device=device, dtype=torch.long)[:, None, None]
|
||||
|
||||
warped_frame.index_put_((batch_indices, trans_pos_floor[:, 1], trans_pos_floor[:, 0]),
|
||||
frame1_cl * weight_nw,
|
||||
accumulate=True)
|
||||
warped_frame.index_put_((batch_indices, trans_pos_ceil[:, 1], trans_pos_floor[:, 0]),
|
||||
frame1_cl * weight_sw,
|
||||
accumulate=True)
|
||||
warped_frame.index_put_((batch_indices, trans_pos_floor[:, 1], trans_pos_ceil[:, 0]),
|
||||
frame1_cl * weight_ne,
|
||||
accumulate=True)
|
||||
warped_frame.index_put_((batch_indices, trans_pos_ceil[:, 1], trans_pos_ceil[:, 0]),
|
||||
frame1_cl * weight_se,
|
||||
accumulate=True)
|
||||
|
||||
warped_weights.index_put_((batch_indices, trans_pos_floor[:, 1], trans_pos_floor[:, 0]), weight_nw, accumulate=True)
|
||||
warped_weights.index_put_((batch_indices, trans_pos_ceil[:, 1], trans_pos_floor[:, 0]), weight_sw, accumulate=True)
|
||||
warped_weights.index_put_((batch_indices, trans_pos_floor[:, 1], trans_pos_ceil[:, 0]), weight_ne, accumulate=True)
|
||||
warped_weights.index_put_((batch_indices, trans_pos_ceil[:, 1], trans_pos_ceil[:, 0]), weight_se, accumulate=True)
|
||||
|
||||
warped_frame_cf = torch.moveaxis(warped_frame, [0, 1, 2, 3], [0, 2, 3, 1])
|
||||
warped_weights_cf = torch.moveaxis(warped_weights, [0, 1, 2, 3], [0, 2, 3, 1])
|
||||
cropped_warped_frame = warped_frame_cf[:, :, 1:-1, 1:-1]
|
||||
cropped_weights = warped_weights_cf[:, :, 1:-1, 1:-1]
|
||||
cropped_weights = torch.nan_to_num(cropped_weights, nan=1000.0)
|
||||
|
||||
mask = cropped_weights > 0
|
||||
zero_value = -1 if is_image else 0
|
||||
zero_tensor = torch.tensor(zero_value, dtype=frame1.dtype, device=frame1.device)
|
||||
warped_frame2 = torch.where(mask, cropped_warped_frame / cropped_weights, zero_tensor)
|
||||
mask2 = mask.to(frame1)
|
||||
|
||||
if is_image:
|
||||
warped_frame2 = torch.clamp(warped_frame2, min=-1, max=1)
|
||||
|
||||
return warped_frame2, mask2
|
||||
|
||||
|
||||
def forward_warp(
|
||||
frame1: torch.Tensor,
|
||||
mask1: torch.Tensor | None,
|
||||
depth1: torch.Tensor | None,
|
||||
transformation1: torch.Tensor | None,
|
||||
transformation2: torch.Tensor,
|
||||
intrinsic1: torch.Tensor | None,
|
||||
intrinsic2: torch.Tensor | None,
|
||||
is_image: bool = True,
|
||||
is_depth: bool = True,
|
||||
render_depth: bool = False,
|
||||
world_points1: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor]:
|
||||
"""
|
||||
Forward warp frame1 to a new view defined by transformation2.
|
||||
|
||||
Args:
|
||||
frame1: (b, c, h, w) source frame in range [-1, 1] for images
|
||||
mask1: (b, 1, h, w) valid pixel mask
|
||||
depth1: (b, 1, h, w) depth map (required if world_points1 is None)
|
||||
transformation1: (b, 4, 4) source camera w2c (required if depth1 is provided)
|
||||
transformation2: (b, 4, 4) target camera w2c
|
||||
intrinsic1: (b, 3, 3) source camera intrinsics
|
||||
intrinsic2: (b, 3, 3) target camera intrinsics
|
||||
is_image: If True, output will be clipped to (-1, 1)
|
||||
is_depth: If True, depth1 is z-depth; if False, it's distance
|
||||
render_depth: If True, also return the warped depth map
|
||||
world_points1: (b, h, w, 3) pre-computed world points (alternative to depth1)
|
||||
|
||||
Returns:
|
||||
warped_frame2: (b, c, h, w) warped frame
|
||||
mask2: (b, 1, h, w) validity mask
|
||||
warped_depth2: (b, h, w) warped depth (if render_depth=True)
|
||||
flow12: (b, 2, h, w) optical flow
|
||||
"""
|
||||
device = frame1.device
|
||||
b, c, h, w = frame1.shape
|
||||
dtype = frame1.dtype
|
||||
|
||||
if mask1 is None:
|
||||
mask1 = torch.ones(size=(b, 1, h, w), device=device, dtype=dtype)
|
||||
if intrinsic2 is None:
|
||||
assert intrinsic1 is not None
|
||||
intrinsic2 = intrinsic1.clone()
|
||||
|
||||
if world_points1 is not None:
|
||||
# Use pre-computed world points
|
||||
assert world_points1.shape == (b, h, w, 3)
|
||||
trans_points1 = project_points(world_points1, transformation2, intrinsic2)
|
||||
else:
|
||||
# Compute from depth
|
||||
assert depth1 is not None and transformation1 is not None
|
||||
assert depth1.shape == (b, 1, h, w)
|
||||
|
||||
depth1 = torch.nan_to_num(depth1, nan=1e4)
|
||||
depth1 = torch.clamp(depth1, min=0, max=1e4)
|
||||
|
||||
# Unproject to world, then project to target view
|
||||
world_points1 = unproject_points(depth1, transformation1, intrinsic1, is_depth=is_depth)
|
||||
trans_points1 = project_points(world_points1, transformation2, intrinsic2)
|
||||
|
||||
# Filter points behind camera
|
||||
mask1 = mask1 * (trans_points1[:, :, :, 2, 0].unsqueeze(1) > 0)
|
||||
trans_coordinates = trans_points1[:, :, :, :2, 0] / (trans_points1[:, :, :, 2:3, 0] + 1e-7)
|
||||
trans_coordinates = trans_coordinates.permute(0, 3, 1, 2) # b, 2, h, w
|
||||
trans_depth1 = trans_points1[:, :, :, 2, 0].unsqueeze(1)
|
||||
|
||||
grid = create_grid(b, h, w, device=device, dtype=dtype)
|
||||
flow12 = trans_coordinates - grid
|
||||
|
||||
warped_frame2, mask2 = bilinear_splatting(frame1, mask1, trans_depth1, flow12, None, is_image=is_image)
|
||||
|
||||
warped_depth2 = None
|
||||
if render_depth:
|
||||
warped_depth2 = bilinear_splatting(trans_depth1, mask1, trans_depth1, flow12, None, is_image=False)[0][:, 0]
|
||||
|
||||
return warped_frame2, mask2, warped_depth2, flow12
|
||||
|
||||
|
||||
def reliable_depth_mask_range_batch(
|
||||
depth: torch.Tensor,
|
||||
window_size: int = 5,
|
||||
ratio_thresh: float = 0.05,
|
||||
eps: float = 1e-6,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute a mask for reliable depth values based on local variation.
|
||||
|
||||
Args:
|
||||
depth: (b, h, w) or (b, 1, h, w) depth map
|
||||
window_size: Size of the local window (must be odd)
|
||||
ratio_thresh: Threshold for depth variation ratio
|
||||
eps: Small epsilon for numerical stability
|
||||
|
||||
Returns:
|
||||
reliable_mask: Boolean mask where True indicates reliable depth
|
||||
"""
|
||||
assert window_size % 2 == 1, "Window size must be odd."
|
||||
|
||||
if depth.dim() == 3:
|
||||
depth_unsq = depth.unsqueeze(1)
|
||||
elif depth.dim() == 4:
|
||||
depth_unsq = depth
|
||||
else:
|
||||
raise ValueError("depth tensor must be of shape (b, h, w) or (b, 1, h, w)")
|
||||
|
||||
local_max = F.max_pool2d(depth_unsq, kernel_size=window_size, stride=1, padding=window_size // 2)
|
||||
local_min = -F.max_pool2d(-depth_unsq, kernel_size=window_size, stride=1, padding=window_size // 2)
|
||||
local_mean = F.avg_pool2d(depth_unsq, kernel_size=window_size, stride=1, padding=window_size // 2)
|
||||
|
||||
ratio = (local_max - local_min) / (local_mean + eps)
|
||||
reliable_mask = (ratio < ratio_thresh) & (depth_unsq > 0)
|
||||
|
||||
return reliable_mask
|
||||
|
||||
|
||||
class Cache3DBase:
|
||||
"""
|
||||
Base class for 3D cache management.
|
||||
|
||||
The cache maintains:
|
||||
- input_image: RGB images stored in the cache
|
||||
- input_points: 3D world coordinates for each pixel
|
||||
- input_mask: Validity mask for each pixel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_image: torch.Tensor,
|
||||
input_depth: torch.Tensor,
|
||||
input_w2c: torch.Tensor,
|
||||
input_intrinsics: torch.Tensor,
|
||||
input_mask: torch.Tensor | None = None,
|
||||
input_format: list[str] | None = None,
|
||||
input_points: torch.Tensor | None = None,
|
||||
weight_dtype: torch.dtype = torch.float32,
|
||||
is_depth: bool = True,
|
||||
device: str = "cuda",
|
||||
filter_points_threshold: float = 1.0,
|
||||
):
|
||||
"""
|
||||
Initialize the 3D cache.
|
||||
|
||||
Args:
|
||||
input_image: Input image tensor with varying dimensions
|
||||
input_depth: Depth map tensor
|
||||
input_w2c: World-to-camera transformation matrix
|
||||
input_intrinsics: Camera intrinsic matrix
|
||||
input_mask: Optional validity mask
|
||||
input_format: Dimension labels for input_image (e.g., ['B', 'C', 'H', 'W'])
|
||||
input_points: Pre-computed 3D world points (alternative to depth)
|
||||
weight_dtype: Data type for computations
|
||||
is_depth: If True, input_depth is z-depth; if False, it's distance
|
||||
device: Computation device
|
||||
filter_points_threshold: Threshold for filtering unreliable depth
|
||||
"""
|
||||
self.weight_dtype = weight_dtype
|
||||
self.is_depth = is_depth
|
||||
self.device = device
|
||||
self.filter_points_threshold = filter_points_threshold
|
||||
|
||||
if input_format is None:
|
||||
assert input_image.dim() == 4
|
||||
input_format = ["B", "C", "H", "W"]
|
||||
|
||||
# Map dimension names to indices
|
||||
format_to_indices = {dim: idx for idx, dim in enumerate(input_format)}
|
||||
input_shape = input_image.shape
|
||||
|
||||
if input_mask is not None:
|
||||
input_image = torch.cat([input_image, input_mask], dim=format_to_indices.get("C"))
|
||||
|
||||
# Extract dimensions
|
||||
B = input_shape[format_to_indices.get("B", 0)] if "B" in format_to_indices else 1
|
||||
F = input_shape[format_to_indices.get("F", 0)] if "F" in format_to_indices else 1
|
||||
N = input_shape[format_to_indices.get("N", 0)] if "N" in format_to_indices else 1
|
||||
V = input_shape[format_to_indices.get("V", 0)] if "V" in format_to_indices else 1
|
||||
H = input_shape[format_to_indices.get("H", 0)] if "H" in format_to_indices else None
|
||||
W = input_shape[format_to_indices.get("W", 0)] if "W" in format_to_indices else None
|
||||
|
||||
# Reorder dimensions to B x F x N x V x C x H x W
|
||||
desired_dims = ["B", "F", "N", "V", "C", "H", "W"]
|
||||
permute_order: list[int | None] = []
|
||||
for dim in desired_dims:
|
||||
idx = format_to_indices.get(dim)
|
||||
permute_order.append(idx)
|
||||
|
||||
permute_indices = [idx for idx in permute_order if idx is not None]
|
||||
input_image = input_image.permute(*permute_indices)
|
||||
|
||||
for i, idx in enumerate(permute_order):
|
||||
if idx is None:
|
||||
input_image = input_image.unsqueeze(i)
|
||||
|
||||
# Now input_image has shape B x F x N x V x C x H x W
|
||||
if input_mask is not None:
|
||||
self.input_image, self.input_mask = input_image[:, :, :, :, :3], input_image[:, :, :, :, 3:]
|
||||
self.input_mask = self.input_mask.to("cpu")
|
||||
else:
|
||||
self.input_mask = None
|
||||
self.input_image = input_image
|
||||
self.input_image = self.input_image.to(weight_dtype).to("cpu")
|
||||
|
||||
# Compute 3D world points
|
||||
if input_points is not None:
|
||||
self.input_points = input_points.reshape(B, F, N, V, H, W, 3).to("cpu")
|
||||
self.input_depth = None
|
||||
else:
|
||||
input_depth = torch.nan_to_num(input_depth, nan=100)
|
||||
input_depth = torch.clamp(input_depth, min=0, max=100)
|
||||
if weight_dtype == torch.float16:
|
||||
input_depth = torch.clamp(input_depth, max=70)
|
||||
|
||||
self.input_points = (unproject_points(
|
||||
input_depth.reshape(-1, 1, H, W),
|
||||
input_w2c.reshape(-1, 4, 4),
|
||||
input_intrinsics.reshape(-1, 3, 3),
|
||||
is_depth=self.is_depth,
|
||||
).to(weight_dtype).reshape(B, F, N, V, H, W, 3).to("cpu"))
|
||||
self.input_depth = input_depth
|
||||
|
||||
# Filter unreliable depth
|
||||
if self.filter_points_threshold < 1.0 and input_depth is not None:
|
||||
input_depth = input_depth.reshape(-1, 1, H, W)
|
||||
depth_mask = reliable_depth_mask_range_batch(input_depth,
|
||||
ratio_thresh=self.filter_points_threshold).reshape(
|
||||
B, F, N, V, 1, H, W)
|
||||
if self.input_mask is None:
|
||||
self.input_mask = depth_mask.to("cpu")
|
||||
else:
|
||||
self.input_mask = self.input_mask * depth_mask.to(self.input_mask.device)
|
||||
|
||||
def update_cache(self, **kwargs):
|
||||
"""Update the cache with new frames. To be implemented by subclasses."""
|
||||
raise NotImplementedError
|
||||
|
||||
def input_frame_count(self) -> int:
|
||||
"""Return the number of frames in the cache."""
|
||||
return self.input_image.shape[1]
|
||||
|
||||
def render_cache(
|
||||
self,
|
||||
target_w2cs: torch.Tensor,
|
||||
target_intrinsics: torch.Tensor,
|
||||
render_depth: bool = False,
|
||||
start_frame_idx: int = 0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Render the cached 3D points from new camera viewpoints.
|
||||
|
||||
Args:
|
||||
target_w2cs: (b, F_target, 4, 4) target camera transformations
|
||||
target_intrinsics: (b, F_target, 3, 3) target camera intrinsics
|
||||
render_depth: If True, return depth instead of RGB
|
||||
start_frame_idx: Starting frame index in the cache
|
||||
|
||||
Returns:
|
||||
pixels: (b, F_target, N, c, h, w) rendered images or depth
|
||||
masks: (b, F_target, N, 1, h, w) validity masks
|
||||
"""
|
||||
bs, F_target, _, _ = target_w2cs.shape
|
||||
B, F, N, V, C, H, W = self.input_image.shape
|
||||
assert bs == B
|
||||
|
||||
target_w2cs = target_w2cs.reshape(B, F_target, 1, 4, 4).expand(B, F_target, N, 4, 4).reshape(-1, 4, 4)
|
||||
target_intrinsics = target_intrinsics.reshape(B, F_target, 1, 3, 3).expand(B, F_target, N, 3,
|
||||
3).reshape(-1, 3, 3)
|
||||
|
||||
# Prepare inputs
|
||||
first_images = rearrange(
|
||||
self.input_image[:, start_frame_idx:start_frame_idx + F_target].expand(B, F_target, N, V, C, H, W),
|
||||
"B F N V C H W -> (B F N) V C H W")
|
||||
first_points = rearrange(
|
||||
self.input_points[:, start_frame_idx:start_frame_idx + F_target].expand(B, F_target, N, V, H, W, 3),
|
||||
"B F N V H W C -> (B F N) V H W C")
|
||||
first_masks = rearrange(
|
||||
self.input_mask[:, start_frame_idx:start_frame_idx + F_target].expand(B, F_target, N, V, 1, H, W),
|
||||
"B F N V C H W -> (B F N) V C H W") if self.input_mask is not None else None
|
||||
|
||||
# Process in chunks for memory efficiency
|
||||
if first_images.shape[1] == 1:
|
||||
warp_chunk_size = 2
|
||||
rendered_warp_images = []
|
||||
rendered_warp_masks = []
|
||||
rendered_warp_depth = []
|
||||
|
||||
first_images = first_images.squeeze(1)
|
||||
first_points = first_points.squeeze(1)
|
||||
first_masks = first_masks.squeeze(1) if first_masks is not None else None
|
||||
|
||||
for i in range(0, first_images.shape[0], warp_chunk_size):
|
||||
with torch.no_grad():
|
||||
imgs_chunk = first_images[i:i + warp_chunk_size].to(self.device, non_blocking=True)
|
||||
pts_chunk = first_points[i:i + warp_chunk_size].to(self.device, non_blocking=True)
|
||||
masks_chunk = (first_masks[i:i + warp_chunk_size].to(self.device, non_blocking=True)
|
||||
if first_masks is not None else None)
|
||||
|
||||
(
|
||||
rendered_warp_images_chunk,
|
||||
rendered_warp_masks_chunk,
|
||||
rendered_warp_depth_chunk,
|
||||
_,
|
||||
) = forward_warp(
|
||||
imgs_chunk,
|
||||
mask1=masks_chunk,
|
||||
depth1=None,
|
||||
transformation1=None,
|
||||
transformation2=target_w2cs[i:i + warp_chunk_size],
|
||||
intrinsic1=target_intrinsics[i:i + warp_chunk_size],
|
||||
intrinsic2=target_intrinsics[i:i + warp_chunk_size],
|
||||
render_depth=render_depth,
|
||||
world_points1=pts_chunk,
|
||||
)
|
||||
|
||||
rendered_warp_images.append(rendered_warp_images_chunk.to("cpu"))
|
||||
rendered_warp_masks.append(rendered_warp_masks_chunk.to("cpu"))
|
||||
if render_depth:
|
||||
rendered_warp_depth.append(rendered_warp_depth_chunk.to("cpu"))
|
||||
|
||||
del imgs_chunk, pts_chunk, masks_chunk
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
rendered_warp_images = torch.cat(rendered_warp_images, dim=0)
|
||||
rendered_warp_masks = torch.cat(rendered_warp_masks, dim=0)
|
||||
if render_depth:
|
||||
rendered_warp_depth = torch.cat(rendered_warp_depth, dim=0)
|
||||
else:
|
||||
raise NotImplementedError("Multi-view rendering not yet supported")
|
||||
|
||||
pixels = rearrange(rendered_warp_images, "(b f n) c h w -> b f n c h w", b=bs, f=F_target, n=N)
|
||||
masks = rearrange(rendered_warp_masks, "(b f n) c h w -> b f n c h w", b=bs, f=F_target, n=N)
|
||||
|
||||
if render_depth:
|
||||
pixels = rearrange(rendered_warp_depth, "(b f n) h w -> b f n h w", b=bs, f=F_target, n=N)
|
||||
|
||||
return pixels.to(self.device), masks.to(self.device)
|
||||
|
||||
|
||||
class Cache3DBuffer(Cache3DBase):
|
||||
"""
|
||||
3D cache with frame buffer support.
|
||||
|
||||
This class manages multiple frame buffers for temporal consistency
|
||||
and supports noise augmentation for training stability.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
frame_buffer_max: int = 2,
|
||||
noise_aug_strength: float = 0.0,
|
||||
generator: torch.Generator | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize the buffered 3D cache.
|
||||
|
||||
Args:
|
||||
frame_buffer_max: Maximum number of frames to buffer
|
||||
noise_aug_strength: Strength of noise augmentation per buffer
|
||||
generator: Random generator for reproducibility
|
||||
**kwargs: Arguments passed to Cache3DBase
|
||||
"""
|
||||
super().__init__(**kwargs)
|
||||
self.frame_buffer_max = frame_buffer_max
|
||||
self.noise_aug_strength = noise_aug_strength
|
||||
self.generator = generator
|
||||
|
||||
def update_cache(
|
||||
self,
|
||||
new_image: torch.Tensor,
|
||||
new_depth: torch.Tensor,
|
||||
new_w2c: torch.Tensor,
|
||||
new_mask: torch.Tensor | None = None,
|
||||
new_intrinsics: torch.Tensor | None = None,
|
||||
):
|
||||
"""
|
||||
Update the cache with a new frame.
|
||||
|
||||
Args:
|
||||
new_image: (B, C, H, W) new RGB image
|
||||
new_depth: (B, 1, H, W) new depth map
|
||||
new_w2c: (B, 4, 4) new world-to-camera transformation
|
||||
new_mask: Optional (B, 1, H, W) validity mask
|
||||
new_intrinsics: (B, 3, 3) camera intrinsics (optional)
|
||||
"""
|
||||
new_image = new_image.to(self.weight_dtype).to(self.device)
|
||||
new_depth = new_depth.to(self.weight_dtype).to(self.device)
|
||||
new_w2c = new_w2c.to(self.weight_dtype).to(self.device)
|
||||
if new_intrinsics is not None:
|
||||
new_intrinsics = new_intrinsics.to(self.weight_dtype).to(self.device)
|
||||
|
||||
new_depth = torch.nan_to_num(new_depth, nan=1e4)
|
||||
new_depth = torch.clamp(new_depth, min=0, max=1e4)
|
||||
|
||||
B, F, N, V, C, H, W = self.input_image.shape
|
||||
|
||||
# Compute new 3D points
|
||||
new_points = unproject_points(new_depth, new_w2c, new_intrinsics, is_depth=self.is_depth).cpu()
|
||||
new_image = new_image.cpu()
|
||||
|
||||
if self.filter_points_threshold < 1.0:
|
||||
new_depth = new_depth.reshape(-1, 1, H, W)
|
||||
depth_mask = reliable_depth_mask_range_batch(new_depth,
|
||||
ratio_thresh=self.filter_points_threshold).reshape(B, 1, H, W)
|
||||
new_mask = depth_mask.to("cpu") if new_mask is None else new_mask * depth_mask.to(new_mask.device)
|
||||
if new_mask is not None:
|
||||
new_mask = new_mask.cpu()
|
||||
|
||||
# Update buffer (newest frame first)
|
||||
if self.frame_buffer_max > 1:
|
||||
if self.input_image.shape[2] < self.frame_buffer_max:
|
||||
self.input_image = torch.cat([new_image[:, None, None, None], self.input_image], 2)
|
||||
self.input_points = torch.cat([new_points[:, None, None, None], self.input_points], 2)
|
||||
if self.input_mask is not None:
|
||||
self.input_mask = torch.cat([new_mask[:, None, None, None], self.input_mask], 2)
|
||||
else:
|
||||
self.input_image[:, :, 0] = new_image[:, None, None]
|
||||
self.input_points[:, :, 0] = new_points[:, None, None]
|
||||
if self.input_mask is not None:
|
||||
self.input_mask[:, :, 0] = new_mask[:, None, None]
|
||||
else:
|
||||
self.input_image = new_image[:, None, None, None]
|
||||
self.input_points = new_points[:, None, None, None]
|
||||
|
||||
def render_cache(
|
||||
self,
|
||||
target_w2cs: torch.Tensor,
|
||||
target_intrinsics: torch.Tensor,
|
||||
render_depth: bool = False,
|
||||
start_frame_idx: int = 0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Render the cache with optional noise augmentation.
|
||||
|
||||
Args:
|
||||
target_w2cs: (b, F_target, 4, 4) target camera transformations
|
||||
target_intrinsics: (b, F_target, 3, 3) target camera intrinsics
|
||||
render_depth: If True, return depth instead of RGB
|
||||
start_frame_idx: Starting frame index (must be 0 for this class)
|
||||
|
||||
Returns:
|
||||
pixels: (b, F_target, N, c, h, w) rendered images
|
||||
masks: (b, F_target, N, 1, h, w) validity masks
|
||||
"""
|
||||
assert start_frame_idx == 0, "start_frame_idx must be 0 for Cache3DBuffer"
|
||||
|
||||
output_device = target_w2cs.device
|
||||
target_w2cs = target_w2cs.to(self.weight_dtype).to(self.device)
|
||||
target_intrinsics = target_intrinsics.to(self.weight_dtype).to(self.device)
|
||||
|
||||
pixels, masks = super().render_cache(target_w2cs, target_intrinsics, render_depth)
|
||||
|
||||
pixels = pixels.to(output_device)
|
||||
masks = masks.to(output_device)
|
||||
|
||||
# Apply noise augmentation (stronger for older buffers)
|
||||
if not render_depth and self.noise_aug_strength > 0:
|
||||
noise = torch.randn(pixels.shape, generator=self.generator, device=pixels.device, dtype=pixels.dtype)
|
||||
per_buffer_noise = (torch.arange(start=pixels.shape[2] - 1, end=-1, step=-1, device=pixels.device) *
|
||||
self.noise_aug_strength)
|
||||
pixels = pixels + noise * per_buffer_noise.reshape(1, 1, -1, 1, 1, 1)
|
||||
|
||||
return pixels, masks
|
||||
@@ -0,0 +1,203 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Ported from NVIDIA GEN3C: cosmos_predict1/diffusion/inference/camera_utils.py
|
||||
"""Camera trajectory generation utilities for GEN3C 3D cache conditioning."""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def apply_transformation(Bx4x4: torch.Tensor, another_matrix: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply batch transformation to a matrix."""
|
||||
B = Bx4x4.shape[0]
|
||||
if another_matrix.dim() == 2:
|
||||
another_matrix = another_matrix.unsqueeze(0).expand(B, -1, -1)
|
||||
return torch.bmm(Bx4x4, another_matrix)
|
||||
|
||||
|
||||
def look_at_matrix(camera_pos: torch.Tensor, target: torch.Tensor, invert_pos: bool = True) -> torch.Tensor:
|
||||
"""Create a 4x4 look-at view matrix pointing camera toward target."""
|
||||
forward = (target - camera_pos).float()
|
||||
forward = forward / torch.norm(forward)
|
||||
|
||||
up = torch.tensor([0.0, 1.0, 0.0], device=camera_pos.device)
|
||||
right = torch.cross(up, forward)
|
||||
right = right / torch.norm(right)
|
||||
up = torch.cross(forward, right)
|
||||
|
||||
look_at = torch.eye(4, device=camera_pos.device)
|
||||
look_at[0, :3] = right
|
||||
look_at[1, :3] = up
|
||||
look_at[2, :3] = forward
|
||||
look_at[:3, 3] = (-camera_pos) if invert_pos else camera_pos
|
||||
|
||||
return look_at
|
||||
|
||||
|
||||
def create_horizontal_trajectory(
|
||||
world_to_camera_matrix: torch.Tensor,
|
||||
center_depth: float,
|
||||
positive: bool = True,
|
||||
n_steps: int = 13,
|
||||
distance: float = 0.1,
|
||||
device: str = "cuda",
|
||||
axis: str = "x",
|
||||
camera_rotation: str = "center_facing",
|
||||
) -> torch.Tensor:
|
||||
"""Create a linear camera trajectory along a specified axis."""
|
||||
look_at_target = torch.tensor([0.0, 0.0, center_depth]).to(device)
|
||||
trajectory = []
|
||||
initial_camera_pos = torch.tensor([0, 0, 0], device=device, dtype=torch.float32)
|
||||
|
||||
translation_positions = []
|
||||
for i in range(n_steps):
|
||||
offset = i * distance * center_depth / n_steps * (1 if positive else -1)
|
||||
if axis == "x":
|
||||
pos = torch.tensor([offset, 0, 0], device=device)
|
||||
elif axis == "y":
|
||||
pos = torch.tensor([0, offset, 0], device=device)
|
||||
elif axis == "z":
|
||||
pos = torch.tensor([0, 0, offset], device=device)
|
||||
else:
|
||||
raise ValueError(f"Axis should be x, y or z, got {axis}")
|
||||
translation_positions.append(pos)
|
||||
|
||||
for pos in translation_positions:
|
||||
camera_pos = initial_camera_pos + pos
|
||||
if camera_rotation == "trajectory_aligned":
|
||||
_look_at = look_at_target + pos * 2
|
||||
elif camera_rotation == "center_facing":
|
||||
_look_at = look_at_target
|
||||
elif camera_rotation == "no_rotation":
|
||||
_look_at = look_at_target + pos
|
||||
else:
|
||||
raise ValueError(f"camera_rotation should be center_facing, trajectory_aligned, "
|
||||
f"or no_rotation, got {camera_rotation}")
|
||||
view_matrix = look_at_matrix(camera_pos, _look_at)
|
||||
trajectory.append(view_matrix)
|
||||
|
||||
trajectory = torch.stack(trajectory)
|
||||
return apply_transformation(trajectory, world_to_camera_matrix)
|
||||
|
||||
|
||||
def create_spiral_trajectory(
|
||||
world_to_camera_matrix: torch.Tensor,
|
||||
center_depth: float,
|
||||
radius_x: float = 0.03,
|
||||
radius_y: float = 0.02,
|
||||
radius_z: float = 0.0,
|
||||
positive: bool = True,
|
||||
camera_rotation: str = "center_facing",
|
||||
n_steps: int = 13,
|
||||
device: str = "cuda",
|
||||
start_from_zero: bool = True,
|
||||
num_circles: int = 1,
|
||||
) -> torch.Tensor:
|
||||
"""Create a spiral/circular camera trajectory."""
|
||||
look_at_target = torch.tensor([0.0, 0.0, center_depth]).to(device)
|
||||
trajectory = []
|
||||
initial_camera_pos = torch.tensor([0, 0, 0], device=device, dtype=torch.float32)
|
||||
|
||||
theta_max = 2 * math.pi * num_circles
|
||||
spiral_positions = []
|
||||
|
||||
for i in range(n_steps):
|
||||
theta = theta_max * i / (n_steps - 1)
|
||||
if start_from_zero:
|
||||
x = radius_x * (math.cos(theta) - 1) * (1 if positive else -1) * center_depth
|
||||
else:
|
||||
x = radius_x * math.cos(theta) * center_depth
|
||||
y = radius_y * math.sin(theta) * center_depth
|
||||
z = radius_z * math.sin(theta) * center_depth
|
||||
spiral_positions.append(torch.tensor([x, y, z], device=device))
|
||||
|
||||
for pos in spiral_positions:
|
||||
camera_pos = initial_camera_pos + pos
|
||||
if camera_rotation == "center_facing":
|
||||
view_matrix = look_at_matrix(camera_pos, look_at_target)
|
||||
elif camera_rotation == "trajectory_aligned":
|
||||
view_matrix = look_at_matrix(camera_pos, look_at_target + pos * 2)
|
||||
elif camera_rotation == "no_rotation":
|
||||
view_matrix = look_at_matrix(camera_pos, look_at_target + pos)
|
||||
else:
|
||||
raise ValueError(f"camera_rotation should be center_facing, trajectory_aligned, "
|
||||
f"or no_rotation, got {camera_rotation}")
|
||||
trajectory.append(view_matrix)
|
||||
|
||||
trajectory = torch.stack(trajectory)
|
||||
return apply_transformation(trajectory, world_to_camera_matrix)
|
||||
|
||||
|
||||
def generate_camera_trajectory(
|
||||
trajectory_type: str,
|
||||
initial_w2c: torch.Tensor,
|
||||
initial_intrinsics: torch.Tensor,
|
||||
num_frames: int,
|
||||
movement_distance: float,
|
||||
camera_rotation: str = "center_facing",
|
||||
center_depth: float = 1.0,
|
||||
device: str = "cuda",
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate camera trajectory for GEN3C video generation.
|
||||
|
||||
Args:
|
||||
trajectory_type: One of "left", "right", "up", "down", "zoom_in",
|
||||
"zoom_out", "clockwise", "counterclockwise".
|
||||
initial_w2c: Initial world-to-camera matrix (4, 4).
|
||||
initial_intrinsics: Camera intrinsics matrix (3, 3).
|
||||
num_frames: Number of frames in the trajectory.
|
||||
movement_distance: Distance factor for camera movement.
|
||||
camera_rotation: "center_facing", "no_rotation", or "trajectory_aligned".
|
||||
center_depth: Depth of the scene center point.
|
||||
device: Computation device.
|
||||
|
||||
Returns:
|
||||
generated_w2cs: (1, num_frames, 4, 4) world-to-camera matrices.
|
||||
generated_intrinsics: (1, num_frames, 3, 3) camera intrinsics.
|
||||
"""
|
||||
if trajectory_type in ["clockwise", "counterclockwise"]:
|
||||
new_w2cs_seq = create_spiral_trajectory(
|
||||
world_to_camera_matrix=initial_w2c,
|
||||
center_depth=center_depth,
|
||||
n_steps=num_frames,
|
||||
positive=trajectory_type == "clockwise",
|
||||
device=device,
|
||||
camera_rotation=camera_rotation,
|
||||
radius_x=movement_distance,
|
||||
radius_y=movement_distance,
|
||||
)
|
||||
elif trajectory_type == "none":
|
||||
# Static camera - repeat identity
|
||||
new_w2cs_seq = initial_w2c.unsqueeze(0).expand(num_frames, -1, -1)
|
||||
else:
|
||||
axis_map = {
|
||||
"left": (False, "x"),
|
||||
"right": (True, "x"),
|
||||
"up": (False, "y"),
|
||||
"down": (True, "y"),
|
||||
"zoom_in": (True, "z"),
|
||||
"zoom_out": (False, "z"),
|
||||
}
|
||||
if trajectory_type not in axis_map:
|
||||
raise ValueError(f"Unsupported trajectory type: {trajectory_type}")
|
||||
positive, axis = axis_map[trajectory_type]
|
||||
|
||||
new_w2cs_seq = create_horizontal_trajectory(
|
||||
world_to_camera_matrix=initial_w2c,
|
||||
center_depth=center_depth,
|
||||
n_steps=num_frames,
|
||||
positive=positive,
|
||||
axis=axis,
|
||||
distance=movement_distance,
|
||||
device=device,
|
||||
camera_rotation=camera_rotation,
|
||||
)
|
||||
|
||||
generated_w2cs = new_w2cs_seq.unsqueeze(0) # (1, num_frames, 4, 4)
|
||||
if initial_intrinsics.dim() == 2:
|
||||
generated_intrinsics = initial_intrinsics.unsqueeze(0).unsqueeze(0).repeat(1, num_frames, 1, 1)
|
||||
else:
|
||||
generated_intrinsics = initial_intrinsics.unsqueeze(0)
|
||||
|
||||
return generated_w2cs, generated_intrinsics
|
||||
@@ -0,0 +1,182 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Ported from NVIDIA GEN3C: cosmos_predict1/diffusion/inference/gen3c_single_image.py
|
||||
"""MoGe-based monocular depth estimation for GEN3C 3D cache conditioning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from moge.model.v1 import MoGeModel
|
||||
else:
|
||||
MoGeModel = Any
|
||||
|
||||
|
||||
def load_moge_model(
|
||||
model_name: str = "Ruicheng/moge-vitl",
|
||||
device: str | torch.device = "cuda",
|
||||
) -> MoGeModel:
|
||||
"""Load MoGe depth estimation model from HuggingFace.
|
||||
|
||||
Args:
|
||||
model_name: HuggingFace model identifier.
|
||||
device: Device to load model on.
|
||||
|
||||
Returns:
|
||||
Loaded MoGe model.
|
||||
"""
|
||||
try:
|
||||
from moge.model.v1 import MoGeModel
|
||||
except ImportError as exc:
|
||||
raise ImportError("MoGe is required for GEN3C 3D cache conditioning. "
|
||||
"Install it with: pip install git+https://github.com/microsoft/MoGe.git. "
|
||||
"If import fails with libGL.so.1, install system deps: "
|
||||
"sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1") from exc
|
||||
|
||||
logger.info("Loading MoGe depth model: %s", model_name)
|
||||
model = MoGeModel.from_pretrained(model_name).to(device)
|
||||
model.eval()
|
||||
logger.info("MoGe model loaded successfully")
|
||||
return model
|
||||
|
||||
|
||||
def predict_depth_from_path(
|
||||
image_path: str,
|
||||
target_h: int,
|
||||
target_w: int,
|
||||
device: torch.device,
|
||||
moge_model: MoGeModel,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict depth, intrinsics, and mask from an image file path.
|
||||
|
||||
Args:
|
||||
image_path: Path to input image (RGB or BGR, any format cv2 supports).
|
||||
target_h: Target height for output tensors.
|
||||
target_w: Target width for output tensors.
|
||||
device: Computation device.
|
||||
moge_model: Loaded MoGe model.
|
||||
|
||||
Returns:
|
||||
image: (1, 1, 3, target_h, target_w) image tensor in [-1, 1].
|
||||
depth: (1, 1, 1, target_h, target_w) depth map.
|
||||
mask: (1, 1, 1, target_h, target_w) confidence mask.
|
||||
w2c: (1, 1, 4, 4) world-to-camera matrix (identity).
|
||||
intrinsics: (1, 1, 3, 3) camera intrinsics.
|
||||
"""
|
||||
import cv2
|
||||
|
||||
input_image_bgr = cv2.imread(image_path)
|
||||
if input_image_bgr is None:
|
||||
raise FileNotFoundError(f"Input image not found: {image_path}")
|
||||
input_image_rgb = cv2.cvtColor(input_image_bgr, cv2.COLOR_BGR2RGB)
|
||||
|
||||
return _predict_depth_core(input_image_rgb, target_h, target_w, device, moge_model)
|
||||
|
||||
|
||||
def predict_depth_from_tensor(
|
||||
image_tensor: torch.Tensor,
|
||||
moge_model: MoGeModel,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict depth and mask from an image tensor (for autoregressive generation).
|
||||
|
||||
Args:
|
||||
image_tensor: (C, H, W) image tensor in [0, 1] range.
|
||||
moge_model: Loaded MoGe model.
|
||||
|
||||
Returns:
|
||||
depth: (1, 1, H, W) depth map.
|
||||
mask: (1, 1, H, W) confidence mask.
|
||||
"""
|
||||
moge_output = moge_model.infer(image_tensor)
|
||||
depth = moge_output["depth"]
|
||||
mask = moge_output["mask"]
|
||||
|
||||
depth = depth.unsqueeze(0).unsqueeze(0)
|
||||
depth = torch.nan_to_num(depth, nan=1e4)
|
||||
depth = torch.clamp(depth, min=0, max=1e4)
|
||||
|
||||
mask = mask.unsqueeze(0).unsqueeze(0)
|
||||
depth = torch.where(mask == 0, torch.tensor(1000.0, device=depth.device), depth)
|
||||
|
||||
return depth, mask
|
||||
|
||||
|
||||
def _predict_depth_core(
|
||||
input_image_rgb: np.ndarray,
|
||||
target_h: int,
|
||||
target_w: int,
|
||||
device: torch.device,
|
||||
moge_model: MoGeModel,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Core depth prediction logic shared between path and tensor inputs."""
|
||||
import cv2
|
||||
|
||||
# MoGe runs at fixed resolution for best results
|
||||
depth_pred_h, depth_pred_w = 720, 1280
|
||||
|
||||
resized = cv2.resize(input_image_rgb, (depth_pred_w, depth_pred_h))
|
||||
img_tensor = torch.tensor(resized / 255.0, dtype=torch.float32, device=device).permute(2, 0, 1)
|
||||
|
||||
# Run MoGe inference
|
||||
moge_output = moge_model.infer(img_tensor)
|
||||
depth_hw = moge_output["depth"]
|
||||
intrinsics_norm = moge_output["intrinsics"]
|
||||
mask_hw = moge_output["mask"]
|
||||
|
||||
# Replace invalid depth with large value
|
||||
depth_hw = torch.where(mask_hw == 0, torch.tensor(1000.0, device=depth_hw.device), depth_hw)
|
||||
|
||||
# Convert normalized intrinsics to pixel coordinates
|
||||
intrinsics_pixel = intrinsics_norm.clone()
|
||||
intrinsics_pixel[0, 0] *= depth_pred_w # fx
|
||||
intrinsics_pixel[1, 1] *= depth_pred_h # fy
|
||||
intrinsics_pixel[0, 2] *= depth_pred_w # cx
|
||||
intrinsics_pixel[1, 2] *= depth_pred_h # cy
|
||||
|
||||
# Scale to target resolution
|
||||
h_scale = target_h / depth_pred_h
|
||||
w_scale = target_w / depth_pred_w
|
||||
|
||||
depth_target = F.interpolate(depth_hw.unsqueeze(0).unsqueeze(0),
|
||||
size=(target_h, target_w),
|
||||
mode='bilinear',
|
||||
align_corners=False).squeeze(0).squeeze(0)
|
||||
|
||||
mask_target = F.interpolate(mask_hw.unsqueeze(0).unsqueeze(0).to(torch.float32),
|
||||
size=(target_h, target_w),
|
||||
mode='nearest').squeeze(0).squeeze(0).to(torch.bool)
|
||||
|
||||
img_target = F.interpolate(img_tensor.unsqueeze(0), size=(target_h, target_w), mode='bilinear',
|
||||
align_corners=False).squeeze(0)
|
||||
|
||||
# Scale intrinsics for target resolution
|
||||
intrinsics_target = intrinsics_pixel.clone()
|
||||
intrinsics_target[0, 0] *= w_scale # fx
|
||||
intrinsics_target[0, 2] *= w_scale # cx
|
||||
intrinsics_target[1, 1] *= h_scale # fy
|
||||
intrinsics_target[1, 2] *= h_scale # cy
|
||||
|
||||
# Format outputs with batch and frame dimensions: (B, F, ...)
|
||||
# Image: [-1, 1] range
|
||||
image_out = (img_target * 2 - 1).unsqueeze(0).unsqueeze(1)
|
||||
|
||||
depth_out = depth_target.unsqueeze(0).unsqueeze(0).unsqueeze(0)
|
||||
depth_out = torch.nan_to_num(depth_out, nan=1e4)
|
||||
depth_out = torch.clamp(depth_out, min=0, max=1e4)
|
||||
|
||||
mask_out = mask_target.unsqueeze(0).unsqueeze(0).unsqueeze(0)
|
||||
|
||||
w2c_out = torch.eye(4, dtype=torch.float32, device=device).unsqueeze(0).unsqueeze(0)
|
||||
intrinsics_out = intrinsics_target.unsqueeze(0).unsqueeze(0)
|
||||
|
||||
return image_out, depth_out, mask_out, w2c_out, intrinsics_out
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GEN3C video diffusion pipeline wiring."""
|
||||
|
||||
from diffusers import EDMEulerScheduler
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (DecodingStage, Gen3CCFGPolicyStage, Gen3CConditioningStage, Gen3CDenoisingStage,
|
||||
Gen3CLatentPreparationStage, InputValidationStage, TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Gen3CPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
GEN3C Video Generation Pipeline.
|
||||
|
||||
This pipeline extends Cosmos with 3D cache support for camera-controlled
|
||||
video generation. When an input image is provided, it runs the full
|
||||
3D cache conditioning pipeline (depth estimation -> point cloud ->
|
||||
camera trajectory -> forward warping -> VAE encoding).
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
scheduler = self.modules.get("scheduler")
|
||||
if scheduler is not None and hasattr(scheduler, "precondition_inputs"):
|
||||
return
|
||||
|
||||
# GEN3C denoising uses EDM preconditioning terms. The converted
|
||||
# model_index may point to FlowMatch scheduler configs that don't
|
||||
# expose precondition_inputs, so force the official EDM scheduler here.
|
||||
logger.warning(
|
||||
"Replacing loaded scheduler (%s) with EDMEulerScheduler for GEN3C parity.",
|
||||
type(scheduler).__name__ if scheduler is not None else "None",
|
||||
)
|
||||
self.modules["scheduler"] = EDMEulerScheduler(
|
||||
sigma_max=80.0,
|
||||
sigma_min=0.0002,
|
||||
sigma_data=float(getattr(fastvideo_args.pipeline_config, "sigma_data", 0.5)),
|
||||
)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="cfg_policy_stage", stage=Gen3CCFGPolicyStage())
|
||||
|
||||
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=Gen3CConditioningStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=Gen3CLatentPreparationStage(scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer"),
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=Gen3CDenoisingStage(transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = Gen3CPipeline
|
||||
@@ -0,0 +1,18 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Sampling profiles for GEN3C models."""
|
||||
from fastvideo.configs.sample.profiles import ModelProfile
|
||||
|
||||
GEN3C_COSMOS_7B = ModelProfile(
|
||||
name="gen3c_cosmos_7b",
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 35,
|
||||
"trajectory_type": "left",
|
||||
"movement_distance": 0.3,
|
||||
"camera_rotation": "center_facing",
|
||||
},
|
||||
)
|
||||
@@ -143,6 +143,11 @@ class ForwardBatch:
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
|
||||
# Camera control inputs (GEN3C)
|
||||
trajectory_type: str | None = None
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# Latent dimensions
|
||||
height_latents: list[int] | int | None = None
|
||||
width_latents: list[int] | int | None = None
|
||||
@@ -170,6 +175,9 @@ class ForwardBatch:
|
||||
eta: float = 0.0
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
# TeaCache
|
||||
enable_teacache: bool = False
|
||||
|
||||
# LTX-2 multi-modal CFG parameters
|
||||
ltx2_cfg_scale_video: float = 1.0
|
||||
ltx2_cfg_scale_audio: float = 1.0
|
||||
|
||||
@@ -32,6 +32,8 @@ from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
|
||||
from fastvideo.pipelines.stages.matrixgame_denoising import (MatrixGameCausalDenoisingStage)
|
||||
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
|
||||
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
|
||||
from fastvideo.pipelines.stages.gen3c_stages import (Gen3CCFGPolicyStage, Gen3CConditioningStage, Gen3CDenoisingStage,
|
||||
Gen3CLatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage, TextEncodingStage)
|
||||
from fastvideo.pipelines.stages.timestep_preparation import (Cosmos25TimestepPreparationStage, TimestepPreparationStage)
|
||||
|
||||
@@ -61,6 +63,10 @@ __all__ = [
|
||||
"MatrixGameCausalDenoisingStage",
|
||||
"HYWorldDenoisingStage",
|
||||
"GameCraftDenoisingStage",
|
||||
"Gen3CCFGPolicyStage",
|
||||
"Gen3CConditioningStage",
|
||||
"Gen3CLatentPreparationStage",
|
||||
"Gen3CDenoisingStage",
|
||||
"CosmosDenoisingStage",
|
||||
"Cosmos25DenoisingStage",
|
||||
"Cosmos25T2WDenoisingStage",
|
||||
|
||||
@@ -61,9 +61,9 @@ class DenoisingStage(PipelineStage):
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
dtype=torch.float16, # TODO(will): hack
|
||||
supported_attention_backends=(AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
|
||||
supported_attention_backends=(AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.BSA_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.SAGE_ATTN_THREE) # hack
|
||||
)
|
||||
|
||||
def forward(
|
||||
|
||||
@@ -0,0 +1,721 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
GEN3C-specific pipeline stages, allowing us to keep conditioning/latent/denoising logic
|
||||
separate from pipeline orchestration.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Gen3CCFGPolicyStage(PipelineStage):
|
||||
"""
|
||||
Explicitly control when GEN3C runs a conditional/unconditional pair (CFG).
|
||||
|
||||
Policies:
|
||||
- legacy: enable CFG only when guidance_scale > 1.0 (current FastVideo behavior)
|
||||
- official_uncond_at_unity: also run CFG at guidance_scale == 1.0
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
policy = getattr(pipeline_config, "cfg_behavior", "legacy")
|
||||
|
||||
if policy == "legacy":
|
||||
batch.do_classifier_free_guidance = batch.guidance_scale > 1.0
|
||||
return batch
|
||||
|
||||
if policy == "official_uncond_at_unity":
|
||||
batch.do_classifier_free_guidance = batch.guidance_scale >= 1.0
|
||||
has_negative_embeds = (batch.negative_prompt_embeds is not None and len(batch.negative_prompt_embeds) > 0)
|
||||
if (batch.do_classifier_free_guidance and batch.negative_prompt is None and not has_negative_embeds):
|
||||
batch.negative_prompt = getattr(pipeline_config, "default_negative_prompt", "")
|
||||
return batch
|
||||
|
||||
raise ValueError(f"Unsupported GEN3C cfg_behavior: {policy}")
|
||||
|
||||
|
||||
class Gen3CConditioningStage(PipelineStage):
|
||||
"""
|
||||
3D cache conditioning stage for GEN3C.
|
||||
|
||||
This stage performs the core GEN3C innovation:
|
||||
1. Loads the input image
|
||||
2. Predicts depth via MoGe
|
||||
3. Initializes a 3D point cloud cache
|
||||
4. Generates a camera trajectory
|
||||
5. Renders warped frames from the cache at each target camera pose
|
||||
6. Stores rendered warps on the batch for VAE encoding in the latent prep stage
|
||||
"""
|
||||
|
||||
def __init__(self, vae=None) -> None:
|
||||
super().__init__()
|
||||
self._moge_model: Any | None = None
|
||||
self._vae = vae
|
||||
|
||||
def _get_moge_model(self, device: torch.device, model_name: str) -> Any:
|
||||
"""Lazy-load MoGe model on first use and ensure it is on target device."""
|
||||
if self._moge_model is None:
|
||||
from fastvideo.pipelines.basic.gen3c.depth_estimation import (load_moge_model)
|
||||
self._moge_model = load_moge_model(model_name, device)
|
||||
else:
|
||||
first_param = next(self._moge_model.parameters(), None)
|
||||
if first_param is not None and first_param.device != device:
|
||||
self._moge_model = self._moge_model.to(device)
|
||||
return self._moge_model
|
||||
|
||||
def _offload_moge(self) -> None:
|
||||
"""Move MoGe to CPU to free GPU memory before denoising."""
|
||||
if self._moge_model is not None and torch.cuda.is_available():
|
||||
self._moge_model = self._moge_model.cpu()
|
||||
torch.cuda.empty_cache()
|
||||
logger.info("MoGe model offloaded to CPU")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Run 3D cache conditioning pipeline."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
device = get_local_torch_device()
|
||||
batch_extra = getattr(batch, "extra", {}) or {}
|
||||
|
||||
image_path = getattr(batch, 'image_path', None) or batch_extra.get("image_path")
|
||||
if image_path is None:
|
||||
logger.info("No image_path provided - skipping 3D cache conditioning "
|
||||
"(will use zero conditioning)")
|
||||
return batch
|
||||
|
||||
logger.info("Running 3D cache conditioning with image: %s", image_path)
|
||||
|
||||
height = getattr(batch, 'height', None) or getattr(pipeline_config, 'video_resolution', (720, 1280))[0]
|
||||
width = getattr(batch, 'width', None) or getattr(pipeline_config, 'video_resolution', (720, 1280))[1]
|
||||
num_frames = getattr(batch, 'num_frames', None) or getattr(pipeline_config, 'num_frames', 121)
|
||||
|
||||
trajectory_type = (getattr(batch, 'trajectory_type', None) or batch_extra.get("trajectory_type")
|
||||
or getattr(pipeline_config, 'default_trajectory_type', 'left'))
|
||||
movement_distance = (getattr(batch, 'movement_distance', None) or batch_extra.get("movement_distance")
|
||||
or getattr(pipeline_config, 'default_movement_distance', 0.3))
|
||||
camera_rotation = (getattr(batch, 'camera_rotation', None) or batch_extra.get("camera_rotation")
|
||||
or getattr(pipeline_config, 'default_camera_rotation', 'center_facing'))
|
||||
|
||||
frame_buffer_max = getattr(pipeline_config, 'frame_buffer_max', 2)
|
||||
noise_aug_strength = getattr(pipeline_config, 'noise_aug_strength', 0.0)
|
||||
filter_points_threshold = getattr(pipeline_config, 'filter_points_threshold', 0.05)
|
||||
|
||||
moge_model_name = getattr(pipeline_config, 'moge_model_name', 'Ruicheng/moge-vitl')
|
||||
|
||||
from fastvideo.pipelines.basic.gen3c.depth_estimation import (predict_depth_from_path)
|
||||
|
||||
moge_model = self._get_moge_model(device, moge_model_name)
|
||||
|
||||
(
|
||||
image_b1chw,
|
||||
depth_b11hw,
|
||||
mask_b11hw,
|
||||
w2c_b144,
|
||||
intrinsics_b133,
|
||||
) = predict_depth_from_path(image_path, height, width, device, moge_model)
|
||||
|
||||
logger.info(
|
||||
"Depth prediction complete. Depth range: [%.3f, %.3f]",
|
||||
depth_b11hw.min().item(),
|
||||
depth_b11hw.max().item(),
|
||||
)
|
||||
|
||||
from fastvideo.pipelines.basic.gen3c.cache_3d import Cache3DBuffer
|
||||
|
||||
seed = getattr(batch, 'seed', None)
|
||||
if seed is None:
|
||||
seed = 42
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
cache = Cache3DBuffer(
|
||||
frame_buffer_max=frame_buffer_max,
|
||||
generator=generator,
|
||||
noise_aug_strength=noise_aug_strength,
|
||||
input_image=image_b1chw[:, 0].clone(),
|
||||
input_depth=depth_b11hw[:, 0],
|
||||
input_w2c=w2c_b144[:, 0],
|
||||
input_intrinsics=intrinsics_b133[:, 0],
|
||||
filter_points_threshold=filter_points_threshold,
|
||||
)
|
||||
|
||||
logger.info("3D cache initialized with %d frame buffer(s)", frame_buffer_max)
|
||||
|
||||
from fastvideo.pipelines.basic.gen3c.camera_utils import (generate_camera_trajectory)
|
||||
|
||||
initial_w2c = w2c_b144[0, 0]
|
||||
initial_intrinsics = intrinsics_b133[0, 0]
|
||||
|
||||
generated_w2cs, generated_intrinsics = generate_camera_trajectory(
|
||||
trajectory_type=trajectory_type,
|
||||
initial_w2c=initial_w2c,
|
||||
initial_intrinsics=initial_intrinsics,
|
||||
num_frames=num_frames,
|
||||
movement_distance=movement_distance,
|
||||
camera_rotation=camera_rotation,
|
||||
center_depth=1.0,
|
||||
device=device.type if isinstance(device, torch.device) else device,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Camera trajectory generated: type=%s, frames=%d, distance=%.3f",
|
||||
trajectory_type,
|
||||
num_frames,
|
||||
movement_distance,
|
||||
)
|
||||
|
||||
rendered_warp_images, rendered_warp_masks = cache.render_cache(
|
||||
generated_w2cs[:, :num_frames],
|
||||
generated_intrinsics[:, :num_frames],
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Cache rendered. Warped images shape: %s, non-zero mask ratio: %.3f",
|
||||
list(rendered_warp_images.shape),
|
||||
(rendered_warp_masks > 0).float().mean().item(),
|
||||
)
|
||||
|
||||
batch.rendered_warp_images = rendered_warp_images.to(device)
|
||||
batch.rendered_warp_masks = rendered_warp_masks.to(device)
|
||||
batch.input_image_conditioning = image_b1chw[:, 0].unsqueeze(2).contiguous().to(device)
|
||||
batch.cache_3d = cache
|
||||
|
||||
if getattr(pipeline_config, "offload_moge_after_depth", True):
|
||||
self._offload_moge()
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class Gen3CLatentPreparationStage(LatentPreparationStage):
|
||||
"""
|
||||
Latent preparation stage for GEN3C.
|
||||
|
||||
This stage prepares latents and encodes 3D cache buffers through the VAE.
|
||||
If rendered warped frames are available on the batch (from Gen3CConditioningStage),
|
||||
they are VAE-encoded to produce real conditioning. Otherwise falls back to zeros.
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler, transformer, vae) -> None:
|
||||
super().__init__(scheduler, transformer)
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Prepare latents and encode 3D cache buffers."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
device = get_local_torch_device()
|
||||
|
||||
if 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
|
||||
|
||||
num_channels_latents = getattr(self.transformer, 'num_channels_latents', 16)
|
||||
|
||||
fallback_num_frames = getattr(pipeline_config, 'num_frames', 121)
|
||||
if not isinstance(fallback_num_frames, int):
|
||||
fallback_num_frames = 121
|
||||
|
||||
num_frames_raw = getattr(batch, 'num_frames', None)
|
||||
if num_frames_raw is None:
|
||||
num_frames_raw = fallback_num_frames
|
||||
if isinstance(num_frames_raw, list):
|
||||
num_frames = int(num_frames_raw[0]) if len(num_frames_raw) > 0 else fallback_num_frames
|
||||
elif isinstance(num_frames_raw, int):
|
||||
num_frames = num_frames_raw
|
||||
else:
|
||||
num_frames = fallback_num_frames
|
||||
if hasattr(self.vae, "get_latent_num_frames"):
|
||||
latent_frames = int(self.vae.get_latent_num_frames(num_frames))
|
||||
else:
|
||||
temporal_ratio = getattr(
|
||||
pipeline_config.vae_config.arch_config,
|
||||
"temporal_compression_ratio",
|
||||
4,
|
||||
)
|
||||
latent_frames = int((num_frames - 1) // temporal_ratio + 1)
|
||||
height = getattr(batch, 'height', 720)
|
||||
width = getattr(batch, 'width', 1280)
|
||||
|
||||
spatial_ratio = getattr(
|
||||
pipeline_config.vae_config.arch_config,
|
||||
"spatial_compression_ratio",
|
||||
8,
|
||||
)
|
||||
latent_height = height // spatial_ratio
|
||||
latent_width = width // spatial_ratio
|
||||
|
||||
generator = getattr(batch, "generator", None)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(f"Expected {batch_size} generators, got {len(generator)}.")
|
||||
|
||||
latents = randn_tensor(
|
||||
(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
),
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
if hasattr(self.scheduler, 'init_noise_sigma'):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
batch.latents = latents
|
||||
batch.batch_size = batch_size
|
||||
batch.height = height
|
||||
batch.width = width
|
||||
batch.latent_height = latent_height
|
||||
batch.latent_width = latent_width
|
||||
batch.latent_frames = latent_frames
|
||||
batch.raw_latent_shape = latents.shape
|
||||
|
||||
frame_buffer_max = getattr(pipeline_config, 'frame_buffer_max', 2)
|
||||
channels_per_buffer = 32
|
||||
buffer_channels = frame_buffer_max * channels_per_buffer
|
||||
|
||||
rendered_warp_images = getattr(batch, 'rendered_warp_images', None)
|
||||
rendered_warp_masks = getattr(batch, 'rendered_warp_masks', None)
|
||||
|
||||
if rendered_warp_images is not None and rendered_warp_masks is not None:
|
||||
logger.info(
|
||||
"Encoding rendered warped frames through VAE (%d buffers)...",
|
||||
rendered_warp_images.shape[2],
|
||||
)
|
||||
|
||||
self.vae = self.vae.to(device)
|
||||
|
||||
if hasattr(self.vae, 'module'):
|
||||
vae_dtype = next(self.vae.module.parameters()).dtype
|
||||
else:
|
||||
vae_dtype = next(self.vae.parameters()).dtype
|
||||
|
||||
condition_video_pose = self.encode_warped_frames(
|
||||
rendered_warp_images,
|
||||
rendered_warp_masks,
|
||||
self.vae,
|
||||
frame_buffer_max,
|
||||
vae_dtype,
|
||||
)
|
||||
batch.condition_video_pose = condition_video_pose.to(device)
|
||||
|
||||
logger.info(
|
||||
"condition_video_pose encoded. Shape: %s, non-zero: %.4f",
|
||||
list(batch.condition_video_pose.shape),
|
||||
(batch.condition_video_pose != 0).float().mean().item(),
|
||||
)
|
||||
|
||||
source_image = getattr(batch, "input_image_conditioning", None)
|
||||
if source_image is None:
|
||||
source_image = rendered_warp_images[:, 0, 0].unsqueeze(2)
|
||||
first_frame = source_image.to(device=device, dtype=vae_dtype)
|
||||
first_latent = self._retrieve_latents(self.vae.encode(first_frame))
|
||||
conditioning_latents = torch.zeros(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=first_latent.dtype,
|
||||
)
|
||||
conditioning_latents[:, :, :first_latent.shape[2], :, :] = first_latent
|
||||
batch.conditioning_latents = conditioning_latents
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
batch.condition_video_input_mask = torch.zeros(
|
||||
batch_size,
|
||||
1,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
batch.condition_video_input_mask[:, :, 0, :, :] = 1.0
|
||||
else:
|
||||
logger.info("No rendered warps available - using zero conditioning")
|
||||
batch.condition_video_pose = torch.zeros(
|
||||
batch_size,
|
||||
buffer_channels,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
batch.condition_video_input_mask = torch.zeros(
|
||||
batch_size,
|
||||
1,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
batch.conditioning_latents = None
|
||||
|
||||
batch.condition_video_augment_sigma = torch.zeros(batch_size, device=device, dtype=torch.float32)
|
||||
batch.cond_indicator = torch.zeros(
|
||||
batch_size,
|
||||
1,
|
||||
latent_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
batch.cond_indicator[:, :, 0, :, :] = 1.0
|
||||
|
||||
ones_padding = torch.ones_like(batch.cond_indicator)
|
||||
zeros_padding = torch.zeros_like(batch.cond_indicator)
|
||||
batch.cond_mask = batch.cond_indicator * ones_padding + (1 - batch.cond_indicator) * zeros_padding
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
batch.uncond_indicator = batch.cond_indicator.clone()
|
||||
batch.uncond_mask = batch.cond_mask.clone()
|
||||
else:
|
||||
batch.uncond_indicator = None
|
||||
batch.uncond_mask = None
|
||||
|
||||
return batch
|
||||
|
||||
@staticmethod
|
||||
def _retrieve_latents(encoder_output: Any) -> torch.Tensor:
|
||||
if hasattr(encoder_output, "latent_dist"):
|
||||
latent_dist = encoder_output.latent_dist
|
||||
if hasattr(latent_dist, "mode"):
|
||||
return latent_dist.mode()
|
||||
if hasattr(latent_dist, "mean"):
|
||||
return latent_dist.mean
|
||||
return latent_dist.sample()
|
||||
if hasattr(encoder_output, "mode"):
|
||||
return encoder_output.mode()
|
||||
if hasattr(encoder_output, "latents"):
|
||||
return encoder_output.latents
|
||||
if hasattr(encoder_output, "sample"):
|
||||
return encoder_output.sample()
|
||||
if isinstance(encoder_output, torch.Tensor):
|
||||
return encoder_output
|
||||
raise AttributeError(f"Unsupported VAE encoder output type: {type(encoder_output)}")
|
||||
|
||||
def encode_warped_frames(
|
||||
self,
|
||||
condition_state: torch.Tensor,
|
||||
condition_state_mask: torch.Tensor,
|
||||
vae: Any,
|
||||
frame_buffer_max: int,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Encode rendered 3D cache buffers through VAE.
|
||||
|
||||
Args:
|
||||
condition_state: (B, T, N, 3, H, W) rendered RGB images in [-1, 1].
|
||||
condition_state_mask: (B, T, N, 1, H, W) rendered masks in [0, 1].
|
||||
vae: VAE encoder.
|
||||
frame_buffer_max: Maximum number of buffers.
|
||||
dtype: Target dtype.
|
||||
|
||||
Returns:
|
||||
latent_condition: (B, buffer_channels, T_latent, H_latent, W_latent)
|
||||
"""
|
||||
assert condition_state.dim() == 6
|
||||
|
||||
condition_state_mask = (condition_state_mask * 2 - 1).repeat(1, 1, 1, 3, 1, 1)
|
||||
|
||||
latent_condition = []
|
||||
num_buffers = condition_state.shape[2]
|
||||
for i in range(num_buffers):
|
||||
img_input = condition_state[:, :, i].permute(0, 2, 1, 3, 4).to(dtype)
|
||||
mask_input = condition_state_mask[:, :, i].permute(0, 2, 1, 3, 4).to(dtype)
|
||||
batched_input = torch.cat([img_input, mask_input], dim=0)
|
||||
batched_latent = self._retrieve_latents(vae.encode(batched_input)).contiguous()
|
||||
current_video_latent, current_mask_latent = batched_latent.chunk(2, dim=0)
|
||||
|
||||
latent_condition.append(current_video_latent)
|
||||
latent_condition.append(current_mask_latent)
|
||||
|
||||
for _ in range(frame_buffer_max - num_buffers):
|
||||
latent_condition.append(torch.zeros_like(current_video_latent))
|
||||
latent_condition.append(torch.zeros_like(current_mask_latent))
|
||||
|
||||
return torch.cat(latent_condition, dim=1)
|
||||
|
||||
|
||||
class Gen3CDenoisingStage(DenoisingStage):
|
||||
"""
|
||||
Denoising stage for GEN3C models.
|
||||
|
||||
This stage extends the base denoising stage with support for:
|
||||
- condition_video_input_mask: Binary mask indicating conditioning frames
|
||||
- condition_video_pose: VAE-encoded 3D cache buffers
|
||||
- condition_video_augment_sigma: Noise augmentation sigma
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None) -> None:
|
||||
super().__init__(transformer, scheduler, pipeline)
|
||||
|
||||
def _has_edm_preconditioning(self) -> bool:
|
||||
return hasattr(self.scheduler, "precondition_inputs")
|
||||
|
||||
def _precondition_inputs(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
sigma: float | torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
if self._has_edm_preconditioning():
|
||||
return self.scheduler.precondition_inputs(sample, sigma)
|
||||
return sample
|
||||
|
||||
@staticmethod
|
||||
def _reverse_precondition_input(
|
||||
xt: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
sigma_data: float,
|
||||
) -> torch.Tensor:
|
||||
c_in = 1.0 / torch.sqrt(sigma**2 + sigma_data**2)
|
||||
return xt / c_in
|
||||
|
||||
@staticmethod
|
||||
def _reverse_precondition_output(
|
||||
latent: torch.Tensor,
|
||||
xt: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
sigma_data: float,
|
||||
) -> torch.Tensor:
|
||||
c_skip = sigma_data**2 / (sigma**2 + sigma_data**2)
|
||||
c_out = sigma * sigma_data / torch.sqrt(sigma**2 + sigma_data**2)
|
||||
return (latent - c_skip * xt) / c_out
|
||||
|
||||
def _augment_noise_with_latent(
|
||||
self,
|
||||
xt: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
latent: torch.Tensor,
|
||||
indicator: torch.Tensor,
|
||||
condition_augment_sigma: float,
|
||||
sigma_data: float,
|
||||
generator: torch.Generator | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
active_indicator = indicator
|
||||
if float(condition_augment_sigma) >= float(sigma.item()):
|
||||
active_indicator = torch.zeros_like(indicator)
|
||||
|
||||
try:
|
||||
noise = torch.randn_like(latent, generator=generator)
|
||||
except TypeError:
|
||||
noise = torch.randn_like(latent)
|
||||
|
||||
augment_sigma = torch.tensor([condition_augment_sigma], device=latent.device, dtype=latent.dtype)
|
||||
augment_latent = latent + noise * augment_sigma
|
||||
augment_latent = self._precondition_inputs(augment_latent, condition_augment_sigma)
|
||||
if self._has_edm_preconditioning():
|
||||
augment_latent_unscaled = self._reverse_precondition_input(augment_latent,
|
||||
sigma=sigma,
|
||||
sigma_data=sigma_data)
|
||||
else:
|
||||
augment_latent_unscaled = augment_latent
|
||||
|
||||
new_xt = active_indicator * augment_latent_unscaled + (1 - active_indicator) * xt
|
||||
return new_xt, latent, active_indicator
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
guidance_scale = batch.guidance_scale
|
||||
fps = getattr(fastvideo_args.pipeline_config, 'fps', 24)
|
||||
sigma_data = float(getattr(fastvideo_args.pipeline_config, "sigma_data", 0.5))
|
||||
condition_augment_sigma = float(getattr(fastvideo_args.pipeline_config, "sigma_conditional", 0.001))
|
||||
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=latents.device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
condition_video_input_mask = getattr(batch, 'condition_video_input_mask', None)
|
||||
condition_video_pose = getattr(batch, 'condition_video_pose', None)
|
||||
condition_video_augment_sigma = getattr(batch, 'condition_video_augment_sigma', None)
|
||||
conditioning_latents = getattr(batch, 'conditioning_latents', None)
|
||||
cond_indicator = getattr(batch, "cond_indicator", None)
|
||||
unconditioning_latents = conditioning_latents
|
||||
uncond_indicator = getattr(batch, "uncond_indicator", None)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
self.scheduler._init_step_index(t)
|
||||
sigma = self.scheduler.sigmas[self.scheduler.step_index].to(device=latents.device, dtype=latents.dtype)
|
||||
|
||||
model_input = latents
|
||||
latent_for_replace = conditioning_latents
|
||||
indicator_for_replace = cond_indicator
|
||||
if (conditioning_latents is not None and cond_indicator is not None):
|
||||
model_input, latent_for_replace, indicator_for_replace = (self._augment_noise_with_latent(
|
||||
latents,
|
||||
sigma=sigma,
|
||||
latent=conditioning_latents,
|
||||
indicator=cond_indicator,
|
||||
condition_augment_sigma=condition_augment_sigma,
|
||||
sigma_data=sigma_data,
|
||||
generator=batch.generator,
|
||||
))
|
||||
|
||||
timestep = t.flatten().expand(latents.size(0))
|
||||
padding_mask = torch.zeros(
|
||||
batch.batch_size,
|
||||
1,
|
||||
batch.height,
|
||||
batch.width,
|
||||
device=model_input.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
model_input_scaled = self.scheduler.scale_model_input(model_input, timestep=t).to(target_dtype)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=model_input_scaled,
|
||||
timestep=timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.prompt_embeds[0].to(target_dtype),
|
||||
fps=fps,
|
||||
condition_video_input_mask=condition_video_input_mask.to(target_dtype)
|
||||
if condition_video_input_mask is not None else None,
|
||||
condition_video_pose=condition_video_pose.to(target_dtype)
|
||||
if condition_video_pose is not None else None,
|
||||
condition_video_augment_sigma=condition_video_augment_sigma
|
||||
if condition_video_augment_sigma is not None else None,
|
||||
padding_mask=padding_mask,
|
||||
)
|
||||
|
||||
if isinstance(noise_pred, tuple):
|
||||
noise_pred = noise_pred[0]
|
||||
cond_pred = noise_pred.float()
|
||||
|
||||
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None:
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
uncond_pose = torch.zeros_like(
|
||||
condition_video_pose) if condition_video_pose is not None else None
|
||||
|
||||
uncond_noise_pred = self.transformer(
|
||||
hidden_states=model_input_scaled,
|
||||
timestep=timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.negative_prompt_embeds[0].to(target_dtype),
|
||||
fps=fps,
|
||||
condition_video_input_mask=condition_video_input_mask.to(target_dtype)
|
||||
if condition_video_input_mask is not None else None,
|
||||
condition_video_pose=uncond_pose.to(target_dtype) if uncond_pose is not None else None,
|
||||
condition_video_augment_sigma=condition_video_augment_sigma
|
||||
if condition_video_augment_sigma is not None else None,
|
||||
padding_mask=padding_mask,
|
||||
)
|
||||
|
||||
if isinstance(uncond_noise_pred, tuple):
|
||||
uncond_noise_pred = uncond_noise_pred[0]
|
||||
uncond_pred = uncond_noise_pred.float()
|
||||
|
||||
pred = cond_pred + guidance_scale * (cond_pred - uncond_pred)
|
||||
else:
|
||||
pred = cond_pred
|
||||
|
||||
model_output = pred
|
||||
if (latent_for_replace is not None and indicator_for_replace is not None):
|
||||
if self._has_edm_preconditioning():
|
||||
latent_unscaled = self._reverse_precondition_output(
|
||||
latent_for_replace,
|
||||
xt=model_input,
|
||||
sigma=sigma,
|
||||
sigma_data=sigma_data,
|
||||
)
|
||||
else:
|
||||
latent_unscaled = latent_for_replace
|
||||
model_output = indicator_for_replace * latent_unscaled + (1 - indicator_for_replace) * model_output
|
||||
|
||||
latents = self.scheduler.step(
|
||||
model_output,
|
||||
t,
|
||||
model_input,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if hasattr(self, "callback_on_step_end") and self.callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in self.callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = self.callback_on_step_end(self, i, t, callback_kwargs)
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -154,7 +154,16 @@ class CudaPlatformBase(Platform):
|
||||
raise ImportError("The Video Sparse Attention backend is not installed. "
|
||||
"To install it, please follow the instructions at: "
|
||||
"https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation ") from e
|
||||
elif selected_backend == AttentionBackendEnum.BSA_ATTN:
|
||||
try:
|
||||
from fastvideo.attention.backends.bsa_attn import ( # noqa: F401
|
||||
BSAAttentionBackend)
|
||||
logger.info("Using BSA Attention backend.")
|
||||
|
||||
return "fastvideo.attention.backends.bsa_attn.BSAAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error("Failed to import BSA Attention backend: %s", str(e))
|
||||
raise ImportError("The BSA Attention backend failed to import.") from e
|
||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||
try:
|
||||
from fastvideo_kernel import moba_attn_varlen # noqa: F401
|
||||
|
||||
@@ -16,6 +16,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
BSA_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
SLA_ATTN = enum.auto()
|
||||
SAGE_SLA_ATTN = enum.auto()
|
||||
|
||||
+149
-20
@@ -19,6 +19,7 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyuan15I2V480PStepDistilledConfig,
|
||||
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig,
|
||||
Hunyuan15SR1080PConfig)
|
||||
@@ -48,9 +49,6 @@ from fastvideo.configs.pipelines.wan import (
|
||||
)
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.cosmos import (
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam, )
|
||||
from fastvideo.configs.sample.cosmos2_5 import Cosmos25SamplingParamBase
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam, HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.hunyuan15 import (Hunyuan15_480P_SamplingParam,
|
||||
Hunyuan15_480P_StepDistilled_I2V_SamplingParam,
|
||||
@@ -135,6 +133,8 @@ class ConfigInfo:
|
||||
|
||||
sampling_param_cls: type[SamplingParam] | None
|
||||
pipeline_config_cls: type[PipelineConfig]
|
||||
workload_types: tuple[WorkloadType, ...]
|
||||
default_profile: str | None = None
|
||||
|
||||
|
||||
# The central registry mapping a model name to its configuration information
|
||||
@@ -150,15 +150,23 @@ _MODEL_NAME_DETECTORS: list[tuple[str, Callable[[str], bool]]] = []
|
||||
def register_configs(
|
||||
sampling_param_cls: type[SamplingParam] | None,
|
||||
pipeline_config_cls: type[PipelineConfig],
|
||||
workload_types: tuple[WorkloadType, ...],
|
||||
hf_model_paths: list[str] | None = None,
|
||||
model_detectors: list[Callable[[str], bool]] | None = None,
|
||||
default_profile: str | None = None,
|
||||
) -> None:
|
||||
"""Register config classes for a model family."""
|
||||
"""Register config classes for a model family.
|
||||
|
||||
workload_types declares which UI workload options this config supports.
|
||||
Use () for configs not exposed as workload options.
|
||||
"""
|
||||
model_id = str(len(_CONFIG_REGISTRY))
|
||||
|
||||
_CONFIG_REGISTRY[model_id] = ConfigInfo(
|
||||
sampling_param_cls=sampling_param_cls,
|
||||
pipeline_config_cls=pipeline_config_cls,
|
||||
workload_types=workload_types,
|
||||
default_profile=default_profile,
|
||||
)
|
||||
|
||||
if hf_model_paths:
|
||||
@@ -231,10 +239,19 @@ def _get_config_info(
|
||||
|
||||
|
||||
def _register_configs() -> None:
|
||||
# Import profile modules so they self-register into the
|
||||
# profile registry. Deferred to here (rather than top-level)
|
||||
# so that fastvideo.registry is sufficiently initialised when
|
||||
# the transitive fastvideo.pipelines.__init__ import fires.
|
||||
import importlib
|
||||
importlib.import_module("fastvideo.pipelines.basic.cosmos.profiles")
|
||||
importlib.import_module("fastvideo.pipelines.basic.gen3c.profiles")
|
||||
|
||||
# LTX-2 (base)
|
||||
register_configs(
|
||||
sampling_param_cls=LTX2BaseSamplingParam,
|
||||
pipeline_config_cls=LTX2T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"Lightricks/LTX-2",
|
||||
"FastVideo/LTX2-base",
|
||||
@@ -248,6 +265,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=LTX2DistilledSamplingParam,
|
||||
pipeline_config_cls=LTX2T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
],
|
||||
@@ -260,6 +278,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_480P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15T2V480PConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
],
|
||||
@@ -275,6 +294,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_480P_StepDistilled_I2V_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15I2V480PStepDistilledConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled",
|
||||
],
|
||||
@@ -282,6 +302,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_720P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15T2V720PConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v",
|
||||
],
|
||||
@@ -289,6 +310,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_720P_Distilled_I2V_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15I2V720PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_i2v_distilled",
|
||||
],
|
||||
@@ -296,6 +318,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_SR_1080P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15SR1080PConfig,
|
||||
workload_types=(),
|
||||
hf_model_paths=["weizhou03/HunyuanVideo-1.5-Diffusers-1080p", "weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR"],
|
||||
)
|
||||
|
||||
@@ -303,6 +326,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=HunyuanSamplingParam,
|
||||
pipeline_config_cls=HunyuanConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo",
|
||||
],
|
||||
@@ -314,6 +338,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=FastHunyuanSamplingParam,
|
||||
pipeline_config_cls=FastHunyuanConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastHunyuan-diffusers",
|
||||
],
|
||||
@@ -323,6 +348,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=HYWorld_SamplingParam,
|
||||
pipeline_config_cls=HYWorldConfig,
|
||||
workload_types=(),
|
||||
hf_model_paths=[
|
||||
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
|
||||
],
|
||||
@@ -333,6 +359,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=HunyuanGameCraftSamplingParam,
|
||||
pipeline_config_cls=HunyuanGameCraftPipelineConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/HunyuanGameCraft-Diffusers",
|
||||
],
|
||||
@@ -342,6 +369,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=LingBotWorld_SamplingParam,
|
||||
pipeline_config_cls=LingBotWorldI2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
],
|
||||
@@ -352,6 +380,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=PipelineConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
],
|
||||
@@ -360,19 +389,34 @@ def _register_configs() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
# LongCat
|
||||
# LongCat (T2V, I2V, VC use same config; workload varies by path)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LongCatT2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
],
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=["FastVideo/LongCat-Video-T2V-Diffusers"],
|
||||
model_detectors=[
|
||||
lambda path: "longcatimagetovideo" in path.lower(),
|
||||
lambda path: "longcatvideocontinuation" in path.lower(),
|
||||
lambda path: "longcat" in path.lower(),
|
||||
lambda path: "longcat" in path.lower() and "i2v" not in path.lower() and "imagetovideo" not in path.lower()
|
||||
and "vc" not in path.lower() and "videocontinuation" not in path.lower(),
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LongCatT2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["FastVideo/LongCat-Video-I2V-Diffusers"],
|
||||
model_detectors=[
|
||||
lambda path: "longcatimagetovideo" in path.lower() or ("longcat" in path.lower() and "i2v" in path.lower()),
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LongCatT2V480PConfig,
|
||||
workload_types=(),
|
||||
hf_model_paths=["FastVideo/LongCat-Video-VC-Diffusers"],
|
||||
model_detectors=[
|
||||
lambda path: "longcatvideocontinuation" in path.lower() or
|
||||
("longcat" in path.lower() and "vc" in path.lower()),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -380,6 +424,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=MatrixGame2_SamplingParam,
|
||||
pipeline_config_cls=MatrixGameI2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
@@ -390,10 +435,25 @@ def _register_configs() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
# GEN3C (must register before generic Cosmos detector)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Gen3CConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/GEN3C-Cosmos-7B-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "gen3c" in path.lower(),
|
||||
],
|
||||
default_profile="gen3c_cosmos_7b",
|
||||
)
|
||||
|
||||
# Cosmos 2.5
|
||||
register_configs(
|
||||
sampling_param_cls=Cosmos25SamplingParamBase,
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Cosmos25Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"KyleShao/Cosmos-Predict2.5-2B-Diffusers",
|
||||
],
|
||||
@@ -404,25 +464,29 @@ def _register_configs() -> None:
|
||||
"cosmos2.5",
|
||||
)),
|
||||
],
|
||||
default_profile="cosmos25_predict2_2b",
|
||||
)
|
||||
|
||||
# Cosmos 2
|
||||
register_configs(
|
||||
sampling_param_cls=Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=CosmosConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "cosmos" in path.lower() and
|
||||
("2.5" not in path.lower() and "2_5" not in path.lower() and "25" not in path.lower()),
|
||||
lambda path: "cosmos" in path.lower() and ("2.5" not in path.lower() and "2_5" not in path.lower() and "25"
|
||||
not in path.lower() and "gen3c" not in path.lower()),
|
||||
],
|
||||
default_profile="cosmos_predict2_2b",
|
||||
)
|
||||
|
||||
# TurboDiffusion
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionT2V_1_3B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionT2V_1_3B_Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
@@ -431,6 +495,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionT2V_14B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionT2V_14B_Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers",
|
||||
],
|
||||
@@ -438,6 +503,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionI2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionI2V_A14B_Config,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
],
|
||||
@@ -447,6 +513,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=WanT2V_1_3B_SamplingParam,
|
||||
pipeline_config_cls=WanT2V480PConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
@@ -455,6 +522,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=WanT2V_14B_SamplingParam,
|
||||
pipeline_config_cls=WanT2V720PConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers",
|
||||
@@ -463,6 +531,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=WanI2V_14B_480P_SamplingParam,
|
||||
pipeline_config_cls=WanI2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
],
|
||||
@@ -471,6 +540,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=WanI2V_14B_720P_SamplingParam,
|
||||
pipeline_config_cls=WanI2V720PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers",
|
||||
],
|
||||
@@ -478,6 +548,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
pipeline_config_cls=WanI2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers",
|
||||
],
|
||||
@@ -485,6 +556,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
pipeline_config_cls=WANV2VConfig,
|
||||
workload_types=(),
|
||||
hf_model_paths=[
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
|
||||
],
|
||||
@@ -492,6 +564,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=FastWanT2V480P_SamplingParam,
|
||||
pipeline_config_cls=FastWan2_1_T2V_480P_Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers",
|
||||
@@ -501,6 +574,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_TI2V_5B_Config,
|
||||
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
],
|
||||
@@ -508,6 +582,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
|
||||
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
|
||||
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
|
||||
@@ -516,6 +591,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_T2V_A14B_Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
],
|
||||
@@ -523,6 +599,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_I2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_I2V_A14B_Config,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
],
|
||||
@@ -530,17 +607,29 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
pipeline_config_cls=SelfForcingWanT2V480PConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "wancausaldmdpipeline" in path.lower()],
|
||||
)
|
||||
# SFWan2.2: T2V and I2V variants by path
|
||||
register_configs(
|
||||
sampling_param_cls=SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
pipeline_config_cls=SelfForcingWan2_2_T2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=["rand0nmr/SFWan2.2-T2V-A14B-Diffusers"],
|
||||
model_detectors=[
|
||||
lambda path: ("sfwan2.2" in path.lower() or "sfwan2_2" in path.lower()) and "i2v" not in path.lower(),
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
pipeline_config_cls=SelfForcingWan2_2_T2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"],
|
||||
model_detectors=[
|
||||
lambda path: ("sfwan2.2" in path.lower() or "sfwan2_2" in path.lower()) and "i2v" in path.lower(),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -548,6 +637,7 @@ def _register_configs() -> None:
|
||||
register_configs(
|
||||
sampling_param_cls=SD35SamplingParam,
|
||||
pipeline_config_cls=SD35Config,
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
hf_model_paths=[
|
||||
"stabilityai/stable-diffusion-3.5-medium",
|
||||
],
|
||||
@@ -635,11 +725,50 @@ def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
|
||||
_register_configs()
|
||||
|
||||
|
||||
def get_registered_model_paths() -> list[str]:
|
||||
"""Return all registered HuggingFace model paths.
|
||||
|
||||
Useful for UIs and tooling that need to enumerate supported models.
|
||||
"""
|
||||
return sorted(_MODEL_HF_PATH_TO_NAME.keys())
|
||||
|
||||
|
||||
def get_registered_models_with_workloads(workload_type: str | None = None, ) -> list[dict[str, Any]]:
|
||||
"""Return models with workload metadata, optionally filtered by workload.
|
||||
|
||||
Args:
|
||||
workload_type: If set (e.g. "t2v", "i2v", "t2i"), only return models
|
||||
that support this workload. If None, return all with workload_types.
|
||||
|
||||
Returns:
|
||||
List of dicts with keys: id, label, workload_types.
|
||||
"""
|
||||
result: list[dict[str, Any]] = []
|
||||
for path in sorted(_MODEL_HF_PATH_TO_NAME.keys()):
|
||||
model_id = _MODEL_HF_PATH_TO_NAME[path]
|
||||
config_info = _CONFIG_REGISTRY.get(model_id)
|
||||
if config_info is None:
|
||||
continue
|
||||
workload_values = [w.value for w in config_info.workload_types]
|
||||
if workload_type is not None and workload_type.lower() not in workload_values:
|
||||
continue
|
||||
label = path.split("/")[-1].replace("-", " ").replace("_", " ")
|
||||
result.append({
|
||||
"id": path,
|
||||
"label": label,
|
||||
"workload_types": workload_values,
|
||||
})
|
||||
return result
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ConfigInfo",
|
||||
"ModelInfo",
|
||||
"get_model_info",
|
||||
"get_pipeline_config_cls_from_name",
|
||||
"get_registered_model_paths",
|
||||
"get_registered_models_with_workloads",
|
||||
"get_sampling_param_cls_for_name",
|
||||
"get_pipeline_config_classes",
|
||||
]
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,502 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.compat import request_to_sampling_param
|
||||
from fastvideo.entrypoints.cli import main as cli_main
|
||||
from fastvideo.entrypoints.cli.generate import GenerateSubcommand
|
||||
from fastvideo.entrypoints.cli.inference_config import (
|
||||
build_generate_run_config,
|
||||
build_serve_config,
|
||||
)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.cli.serve import ServeSubcommand
|
||||
from fastvideo.entrypoints.openai import api_server
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
def _parse_generate_args(argv: list[str]):
|
||||
parser = FlexibleArgumentParser()
|
||||
subparsers = parser.add_subparsers(dest="subparser")
|
||||
GenerateSubcommand().subparser_init(subparsers)
|
||||
args, unknown = parser.parse_known_args(["generate", *argv])
|
||||
args._unknown = unknown
|
||||
return args, unknown
|
||||
|
||||
|
||||
def _parse_serve_args(argv: list[str]):
|
||||
parser = FlexibleArgumentParser()
|
||||
subparsers = parser.add_subparsers(dest="subparser")
|
||||
ServeSubcommand().subparser_init(subparsers)
|
||||
args, unknown = parser.parse_known_args(["serve", *argv])
|
||||
args._unknown = unknown
|
||||
return args, unknown
|
||||
|
||||
|
||||
def test_generate_parser_preserves_unknown_dotted_overrides(tmp_path):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--request.sampling.seed",
|
||||
"42",
|
||||
])
|
||||
|
||||
assert args.config == str(config_path)
|
||||
assert unknown == ["--request.sampling.seed", "42"]
|
||||
|
||||
|
||||
def test_build_generate_run_config_loads_nested_config_and_dotted_overrides(
|
||||
tmp_path,
|
||||
):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
" engine:\n"
|
||||
" num_gpus: 1\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n"
|
||||
" output:\n"
|
||||
" return_frames: true\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--generator.engine.num_gpus",
|
||||
"2",
|
||||
"--request.sampling.seed",
|
||||
"7",
|
||||
])
|
||||
|
||||
config = build_generate_run_config(args, unknown)
|
||||
|
||||
assert config.generator.model_path == "test-model"
|
||||
assert config.generator.engine.num_gpus == 2
|
||||
assert config.request.prompt == "hello"
|
||||
assert config.request.sampling.seed == 7
|
||||
assert config.request.output.return_frames is True
|
||||
|
||||
|
||||
def test_build_generate_run_config_accepts_dashed_dotted_overrides(tmp_path):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--generator.engine.num-gpus",
|
||||
"2",
|
||||
"--request.output.output-path",
|
||||
"outputs/dashed",
|
||||
])
|
||||
|
||||
config = build_generate_run_config(args, unknown)
|
||||
|
||||
assert config.generator.engine.num_gpus == 2
|
||||
assert config.request.output.output_path == "outputs/dashed"
|
||||
|
||||
|
||||
def test_build_generate_run_config_loads_nested_json_config(tmp_path):
|
||||
config_path = tmp_path / "run.json"
|
||||
config_path.write_text(
|
||||
'{"generator":{"model_path":"json-model"},'
|
||||
'"request":{"prompt":"hello"}}',
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args(["--config", str(config_path)])
|
||||
config = build_generate_run_config(args, unknown)
|
||||
|
||||
assert config.generator.model_path == "json-model"
|
||||
assert config.request.prompt == "hello"
|
||||
assert config.request.output.return_frames is False
|
||||
|
||||
|
||||
def test_build_generate_run_config_preserves_model_defaults_for_omitted_request_fields(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(
|
||||
num_frames=81,
|
||||
height=480,
|
||||
width=832,
|
||||
fps=16,
|
||||
guidance_scale=3.0,
|
||||
negative_prompt="model default",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
SamplingParam,
|
||||
"from_pretrained",
|
||||
classmethod(fake_from_pretrained),
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args(["--config", str(config_path)])
|
||||
config = build_generate_run_config(args, unknown)
|
||||
sampling_param = request_to_sampling_param(
|
||||
config.request,
|
||||
model_path=config.generator.model_path,
|
||||
)
|
||||
|
||||
assert sampling_param.num_frames == 81
|
||||
assert sampling_param.height == 480
|
||||
assert sampling_param.width == 832
|
||||
assert sampling_param.fps == 16
|
||||
assert sampling_param.guidance_scale == 3.0
|
||||
assert sampling_param.negative_prompt == "model default"
|
||||
|
||||
|
||||
def test_build_generate_run_config_rejects_flat_config(tmp_path):
|
||||
config_path = tmp_path / "run-flat.yaml"
|
||||
config_path.write_text(
|
||||
"model_path: flat-model\n"
|
||||
"prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args(["--config", str(config_path)])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="top-level 'generator' mapping",
|
||||
):
|
||||
build_generate_run_config(args, unknown)
|
||||
|
||||
|
||||
def test_build_generate_run_config_rejects_non_dotted_overrides(tmp_path):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--num-gpus",
|
||||
"2",
|
||||
])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="CLI overrides must use dotted config paths",
|
||||
):
|
||||
build_generate_run_config(args, unknown)
|
||||
|
||||
|
||||
def test_build_generate_run_config_requires_single_prompt_source(tmp_path):
|
||||
missing_prompt_path = tmp_path / "missing.yaml"
|
||||
missing_prompt_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, unknown = _parse_generate_args(["--config", str(missing_prompt_path)])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="Either request.prompt or request.inputs.prompt_path must be provided",
|
||||
):
|
||||
build_generate_run_config(args, unknown)
|
||||
|
||||
conflicting_prompt_path = tmp_path / "conflict.yaml"
|
||||
conflicting_prompt_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n"
|
||||
" inputs:\n"
|
||||
" prompt_path: prompts.txt\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, unknown = _parse_generate_args(["--config", str(conflicting_prompt_path)])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="Cannot provide both request.prompt and request.inputs.prompt_path",
|
||||
):
|
||||
build_generate_run_config(args, unknown)
|
||||
|
||||
|
||||
def test_build_serve_config_loads_nested_config_and_dotted_overrides(tmp_path):
|
||||
config_path = tmp_path / "serve.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: serve-model\n"
|
||||
"server:\n"
|
||||
" host: 0.0.0.0\n"
|
||||
" port: 8000\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_serve_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--generator.engine.num_gpus",
|
||||
"3",
|
||||
"--server.port",
|
||||
"9100",
|
||||
])
|
||||
|
||||
config = build_serve_config(args, unknown)
|
||||
|
||||
assert config.generator.model_path == "serve-model"
|
||||
assert config.generator.engine.num_gpus == 3
|
||||
assert config.server.host == "0.0.0.0"
|
||||
assert config.server.port == 9100
|
||||
|
||||
|
||||
def test_build_serve_config_rejects_flat_config(tmp_path):
|
||||
config_path = tmp_path / "serve-flat.yaml"
|
||||
config_path.write_text(
|
||||
"model_path: serve-model\n"
|
||||
"host: 127.0.0.1\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_serve_args(["--config", str(config_path)])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="top-level 'generator' mapping",
|
||||
):
|
||||
build_serve_config(args, unknown)
|
||||
|
||||
|
||||
def test_build_serve_config_rejects_non_dotted_overrides(tmp_path):
|
||||
config_path = tmp_path / "serve.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: serve-model\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
args, unknown = _parse_serve_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--port",
|
||||
"9000",
|
||||
])
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="CLI overrides must use dotted config paths",
|
||||
):
|
||||
build_serve_config(args, unknown)
|
||||
|
||||
|
||||
def test_generate_subcommand_requires_config():
|
||||
args, _ = _parse_generate_args([])
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="fastvideo generate requires --config PATH",
|
||||
):
|
||||
GenerateSubcommand().validate(args)
|
||||
|
||||
|
||||
def test_serve_subcommand_requires_config():
|
||||
args, _ = _parse_serve_args([])
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="fastvideo serve requires --config PATH",
|
||||
):
|
||||
ServeSubcommand().validate(args)
|
||||
|
||||
|
||||
def test_generate_subcommand_rejects_non_positive_num_gpus(tmp_path):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello world\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, _ = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--generator.engine.num_gpus",
|
||||
"0",
|
||||
])
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=r"generator\.engine\.num_gpus must be > 0; got 0",
|
||||
):
|
||||
GenerateSubcommand().validate(args)
|
||||
|
||||
|
||||
def test_serve_subcommand_rejects_non_positive_num_gpus(tmp_path):
|
||||
config_path = tmp_path / "serve.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: serve-model\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, _ = _parse_serve_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--generator.engine.num_gpus",
|
||||
"0",
|
||||
])
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=r"generator\.engine\.num_gpus must be > 0; got 0",
|
||||
):
|
||||
ServeSubcommand().validate(args)
|
||||
|
||||
|
||||
def test_generate_subcommand_dispatches_via_typed_config(tmp_path, monkeypatch):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello world\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, _ = _parse_generate_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--request.sampling.num_frames",
|
||||
"81",
|
||||
])
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
class FakeGenerator:
|
||||
|
||||
def generate(self, request):
|
||||
captured["request"] = request
|
||||
return None
|
||||
|
||||
def fake_from_config(cls, config):
|
||||
captured["config"] = config
|
||||
return FakeGenerator()
|
||||
|
||||
monkeypatch.setattr(
|
||||
VideoGenerator,
|
||||
"from_config",
|
||||
classmethod(fake_from_config),
|
||||
)
|
||||
|
||||
GenerateSubcommand().cmd(args)
|
||||
|
||||
request = captured["request"]
|
||||
assert captured["config"].model_path == "test-model"
|
||||
assert request.prompt == "hello world"
|
||||
assert request.sampling.num_frames == 81
|
||||
assert request.output.return_frames is False
|
||||
|
||||
|
||||
def test_serve_subcommand_dispatches_via_typed_config(tmp_path, monkeypatch):
|
||||
config_path = tmp_path / "serve.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: serve-model\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, _ = _parse_serve_args([
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--server.host",
|
||||
"127.0.0.1",
|
||||
"--server.port",
|
||||
"9000",
|
||||
"--server.output_dir",
|
||||
"serve-outputs/",
|
||||
"--generator.engine.num_gpus",
|
||||
"2",
|
||||
])
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_generator_config_to_fastvideo_args(config):
|
||||
captured["config"] = config
|
||||
return SimpleNamespace(model_path=config.model_path)
|
||||
|
||||
def fake_run_server(fastvideo_args, host, port, output_dir):
|
||||
captured["fastvideo_args"] = fastvideo_args
|
||||
captured["host"] = host
|
||||
captured["port"] = port
|
||||
captured["output_dir"] = output_dir
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.entrypoints.cli.serve.generator_config_to_fastvideo_args",
|
||||
fake_generator_config_to_fastvideo_args,
|
||||
)
|
||||
monkeypatch.setattr(api_server, "run_server", fake_run_server)
|
||||
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
assert captured["config"].model_path == "serve-model"
|
||||
assert captured["config"].engine.num_gpus == 2
|
||||
assert captured["host"] == "127.0.0.1"
|
||||
assert captured["port"] == 9000
|
||||
assert captured["output_dir"] == "serve-outputs/"
|
||||
|
||||
|
||||
def test_serve_subcommand_rejects_non_default_default_request(tmp_path):
|
||||
config_path = tmp_path / "serve-default-request.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: serve-model\n"
|
||||
"default_request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
args, _ = _parse_serve_args(["--config", str(config_path)])
|
||||
|
||||
with pytest.raises(
|
||||
NotImplementedError,
|
||||
match="ServeConfig.default_request is not wired",
|
||||
):
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
|
||||
def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch):
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
sys,
|
||||
"argv",
|
||||
["fastvideo", "--config", str(config_path)],
|
||||
)
|
||||
|
||||
with pytest.raises(SystemExit):
|
||||
cli_main.main()
|
||||
@@ -0,0 +1,41 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.api import (
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
config_to_dict,
|
||||
)
|
||||
|
||||
|
||||
def test_run_config_roundtrip_preserves_nested_defaults() -> None:
|
||||
config = RunConfig(
|
||||
generator=GeneratorConfig(model_path="hf://model"),
|
||||
request=GenerationRequest(
|
||||
prompt="hello",
|
||||
sampling=SamplingConfig(num_frames=48, width=832, height=480),
|
||||
),
|
||||
)
|
||||
|
||||
dumped = config_to_dict(config)
|
||||
|
||||
assert dumped["generator"]["model_path"] == "hf://model"
|
||||
assert dumped["generator"]["engine"]["execution_backend"] == "mp"
|
||||
assert dumped["request"]["sampling"]["num_frames"] == 48
|
||||
assert dumped["request"]["sampling"]["guidance_scale_2"] is None
|
||||
assert dumped["request"]["output"]["save_video"] is True
|
||||
|
||||
|
||||
def test_serve_config_includes_server_and_default_request_defaults() -> None:
|
||||
config = ServeConfig(generator=GeneratorConfig(model_path="/models/ltx2"))
|
||||
|
||||
dumped = config_to_dict(config)
|
||||
|
||||
assert dumped["server"] == {
|
||||
"host": "0.0.0.0",
|
||||
"port": 8000,
|
||||
"output_dir": "outputs/",
|
||||
}
|
||||
assert dumped["default_request"]["sampling"]["fps"] == 24
|
||||
assert dumped["default_request"]["runtime"]["enable_teacache"] is False
|
||||
@@ -0,0 +1,94 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import yaml
|
||||
|
||||
from fastvideo.api import (
|
||||
apply_overrides,
|
||||
load_run_config,
|
||||
parse_cli_overrides,
|
||||
)
|
||||
|
||||
|
||||
def test_parse_cli_overrides_casts_supported_scalar_and_collection_types() -> None:
|
||||
parsed = parse_cli_overrides([
|
||||
"--generator.engine.num_gpus",
|
||||
"4",
|
||||
"--request.runtime.enable_teacache=true",
|
||||
"--request.sampling.guidance_scale",
|
||||
"1.5",
|
||||
"--request.prompt",
|
||||
"[\"a\", \"b\"]",
|
||||
"--request.extensions",
|
||||
"{\"ltx2\": {\"initial_latent_path\": \"/tmp/init.pt\"}}",
|
||||
"--request.output.output_video_name",
|
||||
"clip",
|
||||
"--request.state",
|
||||
"null",
|
||||
])
|
||||
|
||||
assert parsed == {
|
||||
"generator.engine.num_gpus": 4,
|
||||
"request.runtime.enable_teacache": True,
|
||||
"request.sampling.guidance_scale": 1.5,
|
||||
"request.prompt": ["a", "b"],
|
||||
"request.extensions": {"ltx2": {"initial_latent_path": "/tmp/init.pt"}},
|
||||
"request.output.output_video_name": "clip",
|
||||
"request.state": None,
|
||||
}
|
||||
|
||||
|
||||
def test_parse_cli_overrides_normalizes_dashed_dotted_keys() -> None:
|
||||
parsed = parse_cli_overrides([
|
||||
"--generator.engine.num-gpus",
|
||||
"2",
|
||||
"--request.output.output-path",
|
||||
"outputs/custom.mp4",
|
||||
])
|
||||
|
||||
assert parsed == {
|
||||
"generator.engine.num_gpus": 2,
|
||||
"request.output.output_path": "outputs/custom.mp4",
|
||||
}
|
||||
|
||||
|
||||
def test_apply_overrides_merges_nested_dicts_without_mutating_source() -> None:
|
||||
original = {
|
||||
"generator": {
|
||||
"model_path": "/models/base",
|
||||
"engine": {"num_gpus": 1},
|
||||
},
|
||||
"request": {},
|
||||
}
|
||||
|
||||
updated = apply_overrides(
|
||||
original,
|
||||
{
|
||||
"generator.engine.num_gpus": 8,
|
||||
"request.extensions.ltx2.initial_latent_path": "/tmp/init.pt",
|
||||
},
|
||||
)
|
||||
|
||||
assert original["generator"]["engine"]["num_gpus"] == 1
|
||||
assert updated["generator"]["engine"]["num_gpus"] == 8
|
||||
assert updated["request"]["extensions"]["ltx2"]["initial_latent_path"] == "/tmp/init.pt"
|
||||
|
||||
|
||||
def test_load_run_config_applies_dotted_overrides_before_validation(tmp_path) -> None:
|
||||
raw = {
|
||||
"generator": {"model_path": "/models/base"},
|
||||
"request": {"prompt": "baseline"},
|
||||
}
|
||||
path = tmp_path / "run.yaml"
|
||||
path.write_text(yaml.safe_dump(raw), encoding="utf-8")
|
||||
|
||||
loaded = load_run_config(path, overrides=[
|
||||
"--generator.pipeline.workload_type",
|
||||
"t2v",
|
||||
"--request.sampling.num_frames",
|
||||
"81",
|
||||
"--request.extensions.ltx2.initial_latent_path",
|
||||
"/tmp/latent.pt",
|
||||
])
|
||||
|
||||
assert loaded.generator.pipeline.workload_type == "t2v"
|
||||
assert loaded.request.sampling.num_frames == 81
|
||||
assert loaded.request.extensions["ltx2"]["initial_latent_path"] == "/tmp/latent.pt"
|
||||
@@ -0,0 +1,202 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api import (
|
||||
config_to_dict,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
load_run_config,
|
||||
load_serve_config,
|
||||
parse_config,
|
||||
PlannedStage,
|
||||
RunConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
|
||||
def test_parse_config_builds_nested_typed_config() -> None:
|
||||
raw = {
|
||||
"generator": {
|
||||
"model_path": "/models/ltx2",
|
||||
"pipeline": {
|
||||
"workload_type": "t2v",
|
||||
"profile": "ltx2_two_stage",
|
||||
},
|
||||
},
|
||||
"request": {
|
||||
"prompt": ["a fox", "a wolf"],
|
||||
"sampling": {
|
||||
"num_frames": 121,
|
||||
"height": 1024,
|
||||
"width": 1536,
|
||||
"guidance_scale": 1.5,
|
||||
},
|
||||
"state": {
|
||||
"kind": "ltx2_continuation",
|
||||
"payload": {"segment_index": 1},
|
||||
},
|
||||
"plan": {
|
||||
"stages": [
|
||||
{
|
||||
"name": "base",
|
||||
"kind": "sample",
|
||||
}
|
||||
]
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
config = parse_config(RunConfig, raw)
|
||||
|
||||
assert config.generator.pipeline.profile == "ltx2_two_stage"
|
||||
assert config.request.prompt == ["a fox", "a wolf"]
|
||||
assert config.request.state == ContinuationState(
|
||||
kind="ltx2_continuation",
|
||||
payload={"segment_index": 1},
|
||||
)
|
||||
assert config.request.plan is not None
|
||||
assert config.request.plan.stages == [PlannedStage(name="base", kind="sample")]
|
||||
|
||||
|
||||
def test_parse_config_accepts_existing_typed_instance() -> None:
|
||||
typed = RunConfig(
|
||||
generator=GeneratorConfig(model_path="/models/base"),
|
||||
request=GenerationRequest(prompt="hello"),
|
||||
)
|
||||
|
||||
assert parse_config(RunConfig, typed) is typed
|
||||
|
||||
|
||||
def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
raw = {
|
||||
"generator": {"model_path": "/models/wan"},
|
||||
"request": {
|
||||
"prompt": "hello",
|
||||
"sampling": {"num_frames": 16},
|
||||
},
|
||||
}
|
||||
path = tmp_path / "run.yaml"
|
||||
path.write_text(yaml.safe_dump(raw), encoding="utf-8")
|
||||
|
||||
loaded = load_run_config(path)
|
||||
|
||||
assert config_to_dict(loaded) == {
|
||||
"generator": {
|
||||
"model_path": "/models/wan",
|
||||
"revision": None,
|
||||
"trust_remote_code": False,
|
||||
"engine": {
|
||||
"num_gpus": 1,
|
||||
"execution_backend": "mp",
|
||||
"parallelism": {
|
||||
"tp_size": -1,
|
||||
"sp_size": -1,
|
||||
"hsdp_replicate_dim": 1,
|
||||
"hsdp_shard_dim": -1,
|
||||
"dist_timeout": None,
|
||||
},
|
||||
"offload": {
|
||||
"dit": True,
|
||||
"dit_layerwise": True,
|
||||
"text_encoder": True,
|
||||
"image_encoder": True,
|
||||
"vae": True,
|
||||
"pin_cpu_memory": True,
|
||||
},
|
||||
"compile": {"enabled": False, "kwargs": {}},
|
||||
"enable_stage_verification": True,
|
||||
"use_fsdp_inference": False,
|
||||
"disable_autocast": False,
|
||||
"quantization": None,
|
||||
},
|
||||
"pipeline": {
|
||||
"workload_type": None,
|
||||
"profile": None,
|
||||
"profile_version": None,
|
||||
"components": {
|
||||
"config_root": None,
|
||||
"pipeline_config_path": None,
|
||||
"text_encoder_weights": None,
|
||||
"transformer_weights": None,
|
||||
"transformer_2_weights": None,
|
||||
"vae_weights": None,
|
||||
"upsampler_weights": None,
|
||||
"lora_path": None,
|
||||
"override_pipeline_cls_name": None,
|
||||
"override_transformer_cls_name": None,
|
||||
},
|
||||
"profile_overrides": {},
|
||||
"experimental": {},
|
||||
},
|
||||
},
|
||||
"request": {
|
||||
"prompt": "hello",
|
||||
"negative_prompt": None,
|
||||
"inputs": {
|
||||
"prompt_path": None,
|
||||
"image_path": None,
|
||||
"video_path": None,
|
||||
"pil_image": None,
|
||||
"pose": None,
|
||||
"mouse_cond": None,
|
||||
"keyboard_cond": None,
|
||||
"grid_sizes": None,
|
||||
"c2ws_plucker_emb": None,
|
||||
"refine_from": None,
|
||||
"stage1_video": None,
|
||||
},
|
||||
"sampling": {
|
||||
"num_videos_per_prompt": 1,
|
||||
"seed": 1024,
|
||||
"num_frames": 16,
|
||||
"height": 720,
|
||||
"width": 1280,
|
||||
"height_sr": 1072,
|
||||
"width_sr": 1920,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 50,
|
||||
"num_inference_steps_sr": 50,
|
||||
"guidance_scale": 1.0,
|
||||
"guidance_scale_2": None,
|
||||
"guidance_rescale": 0.0,
|
||||
"true_cfg_scale": None,
|
||||
"boundary_ratio": None,
|
||||
"sigmas": None,
|
||||
},
|
||||
"runtime": {
|
||||
"enable_teacache": False,
|
||||
"return_trajectory_latents": False,
|
||||
"return_trajectory_decoded": False,
|
||||
},
|
||||
"output": {
|
||||
"output_path": "outputs/",
|
||||
"output_video_name": None,
|
||||
"save_video": True,
|
||||
"return_frames": True,
|
||||
"return_state": False,
|
||||
},
|
||||
"stage_overrides": {},
|
||||
"state": None,
|
||||
"plan": None,
|
||||
"extensions": {},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_load_serve_config_supports_json_roundtrip(tmp_path) -> None:
|
||||
raw = {
|
||||
"generator": {"model_path": "/models/server"},
|
||||
"server": {"port": 9000},
|
||||
"default_request": {"prompt": "serve default"},
|
||||
}
|
||||
path = tmp_path / "serve.json"
|
||||
path.write_text(json.dumps(raw), encoding="utf-8")
|
||||
|
||||
loaded = load_serve_config(path)
|
||||
|
||||
assert isinstance(loaded, ServeConfig)
|
||||
assert loaded.server.port == 9000
|
||||
assert loaded.default_request.prompt == "serve default"
|
||||
@@ -0,0 +1,300 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import importlib
|
||||
import pkgutil
|
||||
import types
|
||||
from pathlib import Path
|
||||
from typing import Any, Union, get_args, get_origin, get_type_hints
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api import RunConfig, ServeConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.cli.generate import GenerateSubcommand
|
||||
from fastvideo.entrypoints.cli.serve import ServeSubcommand
|
||||
from fastvideo.entrypoints.openai import image_api, video_api
|
||||
from fastvideo.entrypoints.openai.protocol import (
|
||||
ImageGenerationsRequest,
|
||||
VideoGenerationsRequest,
|
||||
)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
_INVENTORY_PATH = _REPO_ROOT / "docs" / "design" / "inference_schema_parity_inventory.yaml"
|
||||
|
||||
|
||||
def _load_inventory() -> dict:
|
||||
with open(_INVENTORY_PATH, encoding="utf-8") as f:
|
||||
return yaml.safe_load(f)
|
||||
|
||||
|
||||
def _flatten_status_section(section: dict, valid_statuses: set[str]) -> set[str]:
|
||||
names: set[str] = set()
|
||||
for status, entries in section.items():
|
||||
assert status in valid_statuses, f"Unknown status {status!r} in parity inventory"
|
||||
if isinstance(entries, dict):
|
||||
names.update(entries)
|
||||
elif isinstance(entries, list):
|
||||
names.update(entries)
|
||||
else:
|
||||
raise TypeError(f"Unsupported inventory entry type for {status!r}: {type(entries)!r}")
|
||||
return names
|
||||
|
||||
|
||||
def _get_extra_dataclass_fields(package_name: str, base_cls: type) -> set[str]:
|
||||
package = importlib.import_module(package_name)
|
||||
base_fields = {f.name for f in dataclasses.fields(base_cls)}
|
||||
extras: set[str] = set()
|
||||
for _, modname, _ in pkgutil.iter_modules(package.__path__):
|
||||
if modname == "__pycache__":
|
||||
continue
|
||||
module = importlib.import_module(f"{package_name}.{modname}")
|
||||
for obj in vars(module).values():
|
||||
if (
|
||||
isinstance(obj, type)
|
||||
and dataclasses.is_dataclass(obj)
|
||||
and issubclass(obj, base_cls)
|
||||
and obj is not base_cls
|
||||
):
|
||||
extras.update(f.name for f in dataclasses.fields(obj) if f.name not in base_fields)
|
||||
return extras
|
||||
|
||||
|
||||
def _get_cli_dests(cmd_cls: type) -> set[str]:
|
||||
parser = FlexibleArgumentParser()
|
||||
subparsers = parser.add_subparsers(dest="subparser")
|
||||
command = cmd_cls()
|
||||
subparser = command.subparser_init(subparsers)
|
||||
return {
|
||||
action.dest
|
||||
for action in subparser._actions
|
||||
if action.option_strings and action.dest != "help"
|
||||
}
|
||||
|
||||
|
||||
def _iter_inventory_targets(value: object) -> list[str]:
|
||||
if isinstance(value, str):
|
||||
return _expand_inventory_target(value)
|
||||
if isinstance(value, dict):
|
||||
target = value.get("target")
|
||||
if isinstance(target, str):
|
||||
return _expand_inventory_target(target)
|
||||
return []
|
||||
|
||||
|
||||
def _expand_inventory_target(target: str) -> list[str]:
|
||||
last_dot = target.rfind(".")
|
||||
if last_dot == -1:
|
||||
return [target]
|
||||
prefix = target[:last_dot]
|
||||
leaf = target[last_dot + 1:]
|
||||
if "," not in leaf:
|
||||
return [target]
|
||||
return [f"{prefix}.{part}" for part in leaf.split(",")]
|
||||
|
||||
|
||||
def _config_root_for_target(target: str) -> type:
|
||||
if target.startswith(("generator.", "request.")):
|
||||
return RunConfig
|
||||
if target.startswith(("server.", "default_request.")):
|
||||
return ServeConfig
|
||||
raise AssertionError(f"Unsupported schema target root: {target}")
|
||||
|
||||
|
||||
def _walk_schema_target(root_type: type, target: str) -> None:
|
||||
current_annotation: Any = root_type
|
||||
for depth, segment in enumerate(target.split("."), start=1):
|
||||
current_annotation = _unwrap_schema_annotation(current_annotation)
|
||||
if current_annotation is Any:
|
||||
return
|
||||
origin = get_origin(current_annotation)
|
||||
if origin is dict:
|
||||
return
|
||||
assert dataclasses.is_dataclass(current_annotation), (
|
||||
f"{target!r} diverges at {'.'.join(target.split('.')[:depth - 1]) or '<root>'}: "
|
||||
f"{current_annotation!r} is not a dataclass or open dict boundary"
|
||||
)
|
||||
hints = get_type_hints(current_annotation)
|
||||
assert segment in hints, f"{target!r} missing segment {segment!r}"
|
||||
current_annotation = hints[segment]
|
||||
|
||||
|
||||
def _unwrap_schema_annotation(annotation: Any) -> Any:
|
||||
origin = get_origin(annotation)
|
||||
if origin in {types.UnionType, Union}:
|
||||
args = [arg for arg in get_args(annotation) if arg is not type(None)]
|
||||
if not args:
|
||||
return Any
|
||||
if Any in args:
|
||||
return Any
|
||||
for arg in args:
|
||||
if dataclasses.is_dataclass(arg):
|
||||
return arg
|
||||
if get_origin(arg) is dict:
|
||||
return arg
|
||||
return args[0]
|
||||
return annotation
|
||||
|
||||
|
||||
def test_inventory_file_exists() -> None:
|
||||
assert _INVENTORY_PATH.exists()
|
||||
|
||||
|
||||
def test_inventory_statuses_are_known() -> None:
|
||||
inventory = _load_inventory()
|
||||
valid_statuses = set(inventory["status_definitions"])
|
||||
for section in inventory["surfaces"].values():
|
||||
unknown = set(section) - valid_statuses
|
||||
assert not unknown, f"Unknown statuses in surface inventory: {sorted(unknown)}"
|
||||
|
||||
|
||||
def test_fastvideo_args_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = {f.name for f in dataclasses.fields(FastVideoArgs)}
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["fastvideo_args"],
|
||||
set(inventory["status_definitions"]),
|
||||
)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_pipeline_config_base_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = {f.name for f in dataclasses.fields(PipelineConfig)}
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["pipeline_config_base"],
|
||||
set(inventory["status_definitions"]),
|
||||
)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_pipeline_config_extension_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = _get_extra_dataclass_fields("fastvideo.configs.pipelines", PipelineConfig)
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["pipeline_config_extensions"],
|
||||
set(inventory["status_definitions"]),
|
||||
)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_sampling_param_base_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = {f.name for f in dataclasses.fields(SamplingParam)}
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["sampling_param_base"],
|
||||
set(inventory["status_definitions"]),
|
||||
)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_sampling_param_extension_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = _get_extra_dataclass_fields("fastvideo.configs.sample", SamplingParam)
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["sampling_param_extensions"],
|
||||
set(inventory["status_definitions"]),
|
||||
)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_openai_request_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
valid_statuses = set(inventory["status_definitions"])
|
||||
|
||||
image_expected = set(ImageGenerationsRequest.model_fields)
|
||||
image_actual = _flatten_status_section(
|
||||
inventory["surfaces"]["openai_image_request"],
|
||||
valid_statuses,
|
||||
)
|
||||
assert image_actual == image_expected
|
||||
|
||||
video_expected = set(VideoGenerationsRequest.model_fields)
|
||||
video_actual = _flatten_status_section(
|
||||
inventory["surfaces"]["openai_video_request"],
|
||||
valid_statuses,
|
||||
)
|
||||
assert video_actual == video_expected
|
||||
|
||||
|
||||
def test_cli_dest_inventory_matches_live_parsers() -> None:
|
||||
inventory = _load_inventory()
|
||||
|
||||
generate_expected = set(inventory["cli"]["generate"]["expected_dests"])
|
||||
assert generate_expected == _get_cli_dests(GenerateSubcommand)
|
||||
|
||||
serve_expected = set(inventory["cli"]["serve"]["expected_dests"])
|
||||
assert serve_expected == _get_cli_dests(ServeSubcommand)
|
||||
|
||||
|
||||
def test_review_gap_fields_are_explicitly_inventory_tracked() -> None:
|
||||
inventory = _load_inventory()
|
||||
|
||||
sampling_extensions = inventory["surfaces"]["sampling_param_extensions"]
|
||||
assert "guidance_scale_2" in sampling_extensions["moved"]
|
||||
|
||||
image_request = inventory["surfaces"]["openai_image_request"]
|
||||
video_request = inventory["surfaces"]["openai_video_request"]
|
||||
assert "true_cfg_scale" in image_request["moved"]
|
||||
assert "guidance_scale_2" in video_request["moved"]
|
||||
assert "true_cfg_scale" in video_request["moved"]
|
||||
|
||||
|
||||
def test_inventory_targets_exist_in_typed_schema() -> None:
|
||||
inventory = _load_inventory()
|
||||
target_statuses = {"moved", "profile_owned"}
|
||||
|
||||
for surface in inventory["surfaces"].values():
|
||||
for status, entries in surface.items():
|
||||
if status not in target_statuses or not isinstance(entries, dict):
|
||||
continue
|
||||
for value in entries.values():
|
||||
for target in _iter_inventory_targets(value):
|
||||
if not target.startswith(("generator.", "request.", "server.", "default_request.")):
|
||||
continue
|
||||
_walk_schema_target(_config_root_for_target(target), target)
|
||||
|
||||
|
||||
def test_openai_size_mapping_preserves_width_height_ordering(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
) -> None:
|
||||
inventory = _load_inventory()
|
||||
|
||||
monkeypatch.setattr(image_api, "get_output_dir", lambda: str(tmp_path))
|
||||
image_kwargs = image_api._build_generation_kwargs(
|
||||
request_id="img-test",
|
||||
prompt="test",
|
||||
size="640x360",
|
||||
)
|
||||
assert image_kwargs["width"] == 640
|
||||
assert image_kwargs["height"] == 360
|
||||
|
||||
image_size = inventory["surfaces"]["openai_image_request"]["moved"]["size"]
|
||||
video_size = inventory["surfaces"]["openai_video_request"]["moved"]["size"]
|
||||
assert image_size["target"] == "request.sampling.width,height"
|
||||
assert video_size["target"] == "request.sampling.width,height"
|
||||
|
||||
|
||||
def test_openai_seconds_mapping_preserves_duration_semantics(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
) -> None:
|
||||
inventory = _load_inventory()
|
||||
|
||||
monkeypatch.setattr(video_api, "get_output_dir", lambda: str(tmp_path))
|
||||
request = VideoGenerationsRequest(prompt="test", seconds=4, fps=24)
|
||||
kwargs = video_api._build_generation_kwargs("vid-test", request)
|
||||
assert kwargs["fps"] == 24
|
||||
assert kwargs["num_frames"] == 96
|
||||
|
||||
seconds_entry = inventory["surfaces"]["openai_video_request"][
|
||||
"compatibility_only"
|
||||
]["seconds"]
|
||||
assert seconds_entry["target"] == "request.sampling.num_frames"
|
||||
assert "fps * seconds" in seconds_entry["note"]
|
||||
@@ -0,0 +1,75 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import pytest
|
||||
|
||||
from fastvideo.api import ConfigValidationError, RunConfig, parse_config
|
||||
|
||||
|
||||
def test_unknown_field_error_includes_nested_path() -> None:
|
||||
with pytest.raises(
|
||||
ConfigValidationError,
|
||||
match=r"generator\.engine\.bogus: unknown field",
|
||||
):
|
||||
parse_config(
|
||||
RunConfig,
|
||||
{
|
||||
"generator": {
|
||||
"model_path": "/models/base",
|
||||
"engine": {"bogus": True},
|
||||
},
|
||||
"request": {},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_missing_required_field_error_includes_path() -> None:
|
||||
with pytest.raises(
|
||||
ConfigValidationError,
|
||||
match=r"generator\.model_path: missing required field",
|
||||
):
|
||||
parse_config(
|
||||
RunConfig,
|
||||
{
|
||||
"generator": {},
|
||||
"request": {},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_literal_error_includes_path() -> None:
|
||||
with pytest.raises(
|
||||
ConfigValidationError,
|
||||
match=r"generator\.engine\.execution_backend: expected one of \['mp', 'ray'\]",
|
||||
):
|
||||
parse_config(
|
||||
RunConfig,
|
||||
{
|
||||
"generator": {
|
||||
"model_path": "/models/base",
|
||||
"engine": {"execution_backend": "threaded"},
|
||||
},
|
||||
"request": {},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_nested_type_error_includes_list_path() -> None:
|
||||
with pytest.raises(
|
||||
ConfigValidationError,
|
||||
match=r"request\.plan\.stages\[0\]\.name: expected str",
|
||||
):
|
||||
parse_config(
|
||||
RunConfig,
|
||||
{
|
||||
"generator": {"model_path": "/models/base"},
|
||||
"request": {
|
||||
"plan": {
|
||||
"stages": [
|
||||
{
|
||||
"name": 123,
|
||||
"kind": "sample",
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
},
|
||||
)
|
||||
@@ -1,6 +1,20 @@
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
import warnings
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api import (
|
||||
GenerationRequest,
|
||||
GenerationResult,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
SamplingConfig,
|
||||
load_run_config,
|
||||
)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
|
||||
|
||||
def _new_video_generator() -> VideoGenerator:
|
||||
@@ -8,6 +22,66 @@ def _new_video_generator() -> VideoGenerator:
|
||||
return VideoGenerator.__new__(VideoGenerator)
|
||||
|
||||
|
||||
def _new_runtime_video_generator() -> VideoGenerator:
|
||||
generator = _new_video_generator()
|
||||
generator.fastvideo_args = SimpleNamespace(
|
||||
model_path="test-model",
|
||||
prompt_txt=None,
|
||||
workload_type=SimpleNamespace(value="t2v"),
|
||||
)
|
||||
generator.executor = SimpleNamespace(
|
||||
set_log_queue=lambda queue: None,
|
||||
clear_log_queue=lambda: None,
|
||||
)
|
||||
generator.config = None
|
||||
return generator
|
||||
|
||||
|
||||
def _patch_from_fastvideo_args(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
def fake_from_fastvideo_args(cls, fastvideo_args, *, log_queue=None):
|
||||
generator = cls.__new__(cls)
|
||||
generator.fastvideo_args = fastvideo_args
|
||||
generator.executor = None
|
||||
generator.config = None
|
||||
captured["fastvideo_args"] = fastvideo_args
|
||||
captured["log_queue"] = log_queue
|
||||
return generator
|
||||
|
||||
monkeypatch.setattr(
|
||||
VideoGenerator,
|
||||
"from_fastvideo_args",
|
||||
classmethod(fake_from_fastvideo_args),
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
def _patch_fastvideo_args_from_kwargs(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
def fake_from_kwargs(cls, **kwargs):
|
||||
captured["kwargs"] = kwargs
|
||||
return SimpleNamespace(
|
||||
model_path=kwargs["model_path"],
|
||||
num_gpus=kwargs["num_gpus"],
|
||||
workload_type=WorkloadType.from_string(kwargs.get("workload_type", "t2v")),
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.api.compat.FastVideoArgs.from_kwargs",
|
||||
classmethod(fake_from_kwargs),
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
def _patch_sampling_param_from_pretrained(monkeypatch):
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls()
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
|
||||
|
||||
def test_prepare_output_path_file_sanitization(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
target_dir = tmp_path / "dir"
|
||||
@@ -76,3 +150,420 @@ def test_prepare_output_path_empty_prompt_fallback(tmp_path):
|
||||
assert os.path.dirname(result) == str(out_dir)
|
||||
assert os.path.basename(result) == "output.mp4"
|
||||
|
||||
|
||||
def test_from_config_normalizes_and_translates(monkeypatch):
|
||||
captured = _patch_from_fastvideo_args(monkeypatch)
|
||||
_patch_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(model_path="test-model")
|
||||
config.engine.num_gpus = 2
|
||||
config.pipeline.workload_type = "t2v"
|
||||
|
||||
generator = VideoGenerator.from_config(config)
|
||||
|
||||
assert captured["fastvideo_args"].model_path == "test-model"
|
||||
assert captured["fastvideo_args"].num_gpus == 2
|
||||
assert captured["fastvideo_args"].workload_type.value == "t2v"
|
||||
assert generator.config == config
|
||||
|
||||
|
||||
def test_from_file_loads_generator_from_run_config(tmp_path, monkeypatch):
|
||||
captured = _patch_from_fastvideo_args(monkeypatch)
|
||||
_patch_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
" engine:\n"
|
||||
" num_gpus: 3\n"
|
||||
"request:\n"
|
||||
" prompt: hello\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
VideoGenerator.from_file(str(config_path))
|
||||
|
||||
assert captured["fastvideo_args"].model_path == "test-model"
|
||||
assert captured["fastvideo_args"].num_gpus == 3
|
||||
|
||||
|
||||
def test_from_pretrained_convenience_kwargs_do_not_warn(monkeypatch):
|
||||
captured = _patch_from_fastvideo_args(monkeypatch)
|
||||
fastvideo_args_capture = _patch_fastvideo_args_from_kwargs(monkeypatch)
|
||||
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"test-model",
|
||||
num_gpus=4,
|
||||
use_fsdp_inference=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
)
|
||||
|
||||
assert not caught
|
||||
assert captured["fastvideo_args"].model_path == "test-model"
|
||||
assert captured["fastvideo_args"].num_gpus == 4
|
||||
assert fastvideo_args_capture["kwargs"]["use_fsdp_inference"] is False
|
||||
assert fastvideo_args_capture["kwargs"]["text_encoder_cpu_offload"] is True
|
||||
assert fastvideo_args_capture["kwargs"]["pin_cpu_memory"] is True
|
||||
assert fastvideo_args_capture["kwargs"]["dit_cpu_offload"] is False
|
||||
assert fastvideo_args_capture["kwargs"]["vae_cpu_offload"] is False
|
||||
assert generator.config is not None
|
||||
assert generator.config.model_path == "test-model"
|
||||
assert generator.config.engine.num_gpus == 4
|
||||
|
||||
|
||||
def test_from_pretrained_legacy_only_kwargs_warn(monkeypatch):
|
||||
captured = _patch_from_fastvideo_args(monkeypatch)
|
||||
_patch_fastvideo_args_from_kwargs(monkeypatch)
|
||||
|
||||
with pytest.warns(DeprecationWarning, match="legacy-only kwargs"):
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"test-model",
|
||||
num_gpus=4,
|
||||
workload_type="t2v",
|
||||
)
|
||||
|
||||
assert captured["fastvideo_args"].model_path == "test-model"
|
||||
assert captured["fastvideo_args"].num_gpus == 4
|
||||
assert captured["fastvideo_args"].workload_type.value == "t2v"
|
||||
assert generator.config is not None
|
||||
assert generator.config.pipeline.workload_type == "t2v"
|
||||
|
||||
|
||||
def test_generate_uses_typed_request_path(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
_patch_sampling_param_from_pretrained(monkeypatch)
|
||||
captured = {}
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["prompt"] = prompt
|
||||
captured["sampling_param"] = sampling_param
|
||||
captured["kwargs"] = kwargs
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
result = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt="hello world",
|
||||
sampling=SamplingConfig(num_frames=81, height=480, width=832),
|
||||
)
|
||||
)
|
||||
|
||||
assert isinstance(result, GenerationResult)
|
||||
assert captured["prompt"] == "hello world"
|
||||
assert captured["sampling_param"].num_frames == 81
|
||||
assert captured["sampling_param"].height == 480
|
||||
assert captured["sampling_param"].width == 832
|
||||
assert result.video_path == "outputs/test.mp4"
|
||||
|
||||
|
||||
def test_generate_preserves_schema_defaults_for_dataclass_request(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(
|
||||
negative_prompt="model default",
|
||||
num_frames=61,
|
||||
height=448,
|
||||
width=832,
|
||||
)
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["sampling_param"] = sampling_param
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt="hello world",
|
||||
negative_prompt=None,
|
||||
sampling=SamplingConfig(num_frames=125, height=720, width=1280),
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["sampling_param"].negative_prompt is None
|
||||
assert captured["sampling_param"].num_frames == 125
|
||||
assert captured["sampling_param"].height == 720
|
||||
assert captured["sampling_param"].width == 1280
|
||||
|
||||
|
||||
def test_generate_mapping_request_preserves_model_defaults_for_omitted_fields(
|
||||
monkeypatch,
|
||||
):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(
|
||||
negative_prompt="model default",
|
||||
num_frames=61,
|
||||
height=448,
|
||||
width=832,
|
||||
fps=16,
|
||||
guidance_scale=3.0,
|
||||
)
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["sampling_param"] = sampling_param
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
generator.generate(
|
||||
{
|
||||
"prompt": "hello world",
|
||||
}
|
||||
)
|
||||
|
||||
assert captured["sampling_param"].negative_prompt == "model default"
|
||||
assert captured["sampling_param"].num_frames == 61
|
||||
assert captured["sampling_param"].height == 448
|
||||
assert captured["sampling_param"].width == 832
|
||||
assert captured["sampling_param"].fps == 16
|
||||
assert captured["sampling_param"].guidance_scale == 3.0
|
||||
|
||||
|
||||
def test_generate_honors_post_load_request_mutations(monkeypatch, tmp_path):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello world\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(seed=1024, num_frames=61, height=448, width=832)
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["sampling_param"] = sampling_param
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
config = load_run_config(config_path)
|
||||
config.request.sampling.seed = 7
|
||||
|
||||
generator.generate(config.request)
|
||||
|
||||
assert captured["sampling_param"].seed == 7
|
||||
|
||||
|
||||
def test_generate_honors_post_load_mutations_matching_schema_defaults(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello world\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(guidance_scale=3.0)
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["sampling_param"] = sampling_param
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
config = load_run_config(config_path)
|
||||
config.request.sampling.guidance_scale = 1.0
|
||||
|
||||
generator.generate(config.request)
|
||||
|
||||
assert captured["sampling_param"].guidance_scale == 1.0
|
||||
|
||||
|
||||
def test_generate_removes_deleted_loaded_stage_overrides(monkeypatch, tmp_path):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
config_path = tmp_path / "run.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
" model_path: test-model\n"
|
||||
"request:\n"
|
||||
" prompt: hello world\n"
|
||||
" stage_overrides:\n"
|
||||
" refine:\n"
|
||||
" t_thresh: 0.8\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
def fake_from_pretrained(cls, model_path):
|
||||
return cls(t_thresh=0.5)
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured["sampling_param"] = sampling_param
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(SamplingParam, "from_pretrained", classmethod(fake_from_pretrained))
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
config = load_run_config(config_path)
|
||||
del config.request.stage_overrides["refine"]
|
||||
|
||||
generator.generate(config.request)
|
||||
|
||||
assert captured["sampling_param"].t_thresh == 0.5
|
||||
|
||||
|
||||
def test_generate_video_legacy_call_uses_legacy_impl(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
|
||||
def fake_generate_video_impl(
|
||||
prompt=None,
|
||||
sampling_param=None,
|
||||
mouse_cond=None,
|
||||
keyboard_cond=None,
|
||||
grid_sizes=None,
|
||||
**kwargs,
|
||||
):
|
||||
captured["prompt"] = prompt
|
||||
captured["sampling_param"] = sampling_param
|
||||
captured["mouse_cond"] = mouse_cond
|
||||
captured["keyboard_cond"] = keyboard_cond
|
||||
captured["grid_sizes"] = grid_sizes
|
||||
captured["kwargs"] = kwargs
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
with pytest.warns(DeprecationWarning):
|
||||
result = generator.generate_video(
|
||||
prompt="legacy prompt",
|
||||
num_frames=49,
|
||||
output_path="outputs/legacy",
|
||||
save_video=False,
|
||||
log_queue="queue-token",
|
||||
)
|
||||
|
||||
assert captured["prompt"] == "legacy prompt"
|
||||
assert captured["kwargs"]["num_frames"] == 49
|
||||
assert captured["kwargs"]["output_path"] == "outputs/legacy"
|
||||
assert captured["kwargs"]["save_video"] is False
|
||||
assert result["video_path"] == "outputs/test.mp4"
|
||||
|
||||
|
||||
def test_generate_video_legacy_call_preserves_unknown_kwargs(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
captured = {}
|
||||
|
||||
def fake_generate_video_impl(
|
||||
prompt=None,
|
||||
sampling_param=None,
|
||||
mouse_cond=None,
|
||||
keyboard_cond=None,
|
||||
grid_sizes=None,
|
||||
**kwargs,
|
||||
):
|
||||
captured["prompt"] = prompt
|
||||
captured["sampling_param"] = sampling_param
|
||||
captured["kwargs"] = kwargs
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
with pytest.warns(DeprecationWarning):
|
||||
result = generator.generate_video(
|
||||
prompt="legacy prompt",
|
||||
neg_prompt="custom negative",
|
||||
embedded_cfg_scale=7.5,
|
||||
)
|
||||
|
||||
assert captured["prompt"] == "legacy prompt"
|
||||
assert captured["kwargs"]["neg_prompt"] == "custom negative"
|
||||
assert captured["kwargs"]["embedded_cfg_scale"] == 7.5
|
||||
assert result["video_path"] == "outputs/test.mp4"
|
||||
|
||||
|
||||
def test_generate_batch_prompt_file_returns_typed_results(tmp_path, monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
_patch_sampling_param_from_pretrained(monkeypatch)
|
||||
prompt_file = tmp_path / "prompts.txt"
|
||||
prompt_file.write_text("first prompt\nsecond prompt\n", encoding="utf-8")
|
||||
output_dir = tmp_path / "outputs"
|
||||
captured_prompts = []
|
||||
|
||||
def fake_generate_single_video(prompt, sampling_param=None, **kwargs):
|
||||
captured_prompts.append(prompt)
|
||||
return {"prompts": prompt, "video_path": kwargs["output_path"]}
|
||||
|
||||
monkeypatch.setattr(generator, "_generate_single_video", fake_generate_single_video)
|
||||
|
||||
results = generator.generate(
|
||||
{
|
||||
"inputs": {"prompt_path": str(prompt_file)},
|
||||
"output": {
|
||||
"output_path": str(output_dir),
|
||||
"save_video": False,
|
||||
"return_frames": False,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert isinstance(results, list)
|
||||
assert [result.prompt for result in results] == ["first prompt", "second prompt"]
|
||||
assert [result.prompt_index for result in results] == [0, 1]
|
||||
assert captured_prompts == ["first prompt", "second prompt"]
|
||||
|
||||
|
||||
def test_generate_batched_request_fans_out_media_inputs(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
_patch_sampling_param_from_pretrained(monkeypatch)
|
||||
captured: list[tuple[str | None, str | None, str | None]] = []
|
||||
|
||||
def fake_generate_video_impl(prompt=None, sampling_param=None, **kwargs):
|
||||
captured.append((prompt, sampling_param.image_path, sampling_param.video_path))
|
||||
return {"prompts": prompt, "video_path": "outputs/test.mp4"}
|
||||
|
||||
monkeypatch.setattr(generator, "_generate_video_impl", fake_generate_video_impl)
|
||||
|
||||
results = generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=["first prompt", "second prompt"],
|
||||
inputs=InputConfig(
|
||||
image_path=["first.png", "second.png"],
|
||||
video_path=["first.mp4", "second.mp4"],
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert [result.prompt for result in results] == ["first prompt", "second prompt"]
|
||||
assert captured == [
|
||||
("first prompt", "first.png", "first.mp4"),
|
||||
("second prompt", "second.png", "second.mp4"),
|
||||
]
|
||||
|
||||
|
||||
def test_generate_batched_request_rejects_mismatched_media_inputs(monkeypatch):
|
||||
generator = _new_runtime_video_generator()
|
||||
_patch_sampling_param_from_pretrained(monkeypatch)
|
||||
|
||||
with pytest.raises(ValueError, match="image_path"):
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=["first prompt", "second prompt"],
|
||||
inputs=InputConfig(image_path=["first.png"]),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def test_inference_bsa():
|
||||
"""Test FastVideo BSA_ATTN inference pipeline"""
|
||||
|
||||
output_dir = Path("outputs_video/bsa_1.3B/")
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "BSA_ATTN"
|
||||
|
||||
config = {
|
||||
"generator": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"engine": {
|
||||
"num_gpus": 1,
|
||||
"parallelism": {
|
||||
"tp_size": 1,
|
||||
"sp_size": 1
|
||||
},
|
||||
"offload": {
|
||||
"dit": False,
|
||||
"vae": False,
|
||||
"text_encoder": True,
|
||||
"pin_cpu_memory": False,
|
||||
},
|
||||
},
|
||||
"pipeline": {
|
||||
"experimental": {
|
||||
"flow_shift": 8.0,
|
||||
},
|
||||
},
|
||||
},
|
||||
"request": {
|
||||
"prompt":
|
||||
"A majestic lion strides across the golden savanna, "
|
||||
"its powerful frame glistening under the warm afternoon "
|
||||
"sun. The tall grass ripples gently in the breeze, "
|
||||
"enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low "
|
||||
"angle, steady tracking shot, cinematic.",
|
||||
"negative_prompt":
|
||||
"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",
|
||||
"sampling": {
|
||||
"seed": 1024,
|
||||
"num_frames": 77,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 10,
|
||||
"guidance_scale": 6.0,
|
||||
},
|
||||
"output": {
|
||||
"output_path": str(output_dir),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
config_path = f.name
|
||||
|
||||
try:
|
||||
cmd = [
|
||||
sys.executable, "-m", "fastvideo.entrypoints.cli.main",
|
||||
"generate", "--config", config_path
|
||||
]
|
||||
subprocess.run(cmd, check=True)
|
||||
finally:
|
||||
os.unlink(config_path)
|
||||
|
||||
assert output_dir.exists(), \
|
||||
f"Output directory {output_dir} does not exist"
|
||||
|
||||
video_files = list(output_dir.glob("*.mp4"))
|
||||
assert len(video_files) > 0, "No video files were generated"
|
||||
|
||||
for video_file in video_files:
|
||||
assert video_file.stat().st_size > 0, \
|
||||
f"Video file {video_file} is empty"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_inference_bsa()
|
||||
@@ -1,58 +1,96 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def test_inference_vmoba():
|
||||
"""Test FastVideo VMOBA_ATTN inference pipeline"""
|
||||
|
||||
num_gpus = "1"
|
||||
model_base = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
output_dir = Path("outputs_video/vmoba_1.3B/")
|
||||
moba_config = "fastvideo/configs/backend/vmoba/wan_1.3B_77_480_832.json"
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VMOBA_ATTN"
|
||||
|
||||
cmd = [
|
||||
"fastvideo", "generate",
|
||||
"--model-path", model_base,
|
||||
"--sp-size", num_gpus,
|
||||
"--tp-size", "1",
|
||||
"--num-gpus", num_gpus,
|
||||
"--dit-cpu-offload", "False",
|
||||
"--vae-cpu-offload", "False",
|
||||
"--text-encoder-cpu-offload", "True",
|
||||
"--pin-cpu-memory", "False",
|
||||
"--height", "480",
|
||||
"--width", "832",
|
||||
"--num-frames", "77",
|
||||
"--num-inference-steps", "10",
|
||||
"--moba-config-path", moba_config,
|
||||
"--fps", "16",
|
||||
"--guidance-scale", "6.0",
|
||||
"--flow-shift", "8.0",
|
||||
"--prompt", "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.",
|
||||
"--negative-prompt", (
|
||||
"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"
|
||||
),
|
||||
"--seed", "1024",
|
||||
"--output-path", str(output_dir),
|
||||
]
|
||||
config = {
|
||||
"generator": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"engine": {
|
||||
"num_gpus": 1,
|
||||
"parallelism": {
|
||||
"tp_size": 1,
|
||||
"sp_size": 1
|
||||
},
|
||||
"offload": {
|
||||
"dit": False,
|
||||
"vae": False,
|
||||
"text_encoder": True,
|
||||
"pin_cpu_memory": False,
|
||||
},
|
||||
},
|
||||
"pipeline": {
|
||||
"experimental": {
|
||||
"flow_shift": 8.0,
|
||||
"moba_config_path": moba_config,
|
||||
},
|
||||
},
|
||||
},
|
||||
"request": {
|
||||
"prompt":
|
||||
"A majestic lion strides across the golden savanna, "
|
||||
"its powerful frame glistening under the warm afternoon "
|
||||
"sun. The tall grass ripples gently in the breeze, "
|
||||
"enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low "
|
||||
"angle, steady tracking shot, cinematic.",
|
||||
"negative_prompt":
|
||||
"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",
|
||||
"sampling": {
|
||||
"seed": 1024,
|
||||
"num_frames": 77,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 10,
|
||||
"guidance_scale": 6.0,
|
||||
},
|
||||
"output": {
|
||||
"output_path": str(output_dir),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
subprocess.run(cmd, check=True)
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
config_path = f.name
|
||||
|
||||
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
|
||||
try:
|
||||
cmd = ["fastvideo", "generate", "--config", config_path]
|
||||
subprocess.run(cmd, check=True)
|
||||
finally:
|
||||
os.unlink(config_path)
|
||||
|
||||
assert output_dir.exists(), \
|
||||
f"Output directory {output_dir} does not exist"
|
||||
|
||||
video_files = list(output_dir.glob("*.mp4"))
|
||||
assert len(video_files) > 0, "No video files were generated"
|
||||
|
||||
for video_file in video_files:
|
||||
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
|
||||
assert video_file.stat().st_size > 0, \
|
||||
f"Video file {video_file} is empty"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_inference_vmoba()
|
||||
|
||||
@@ -206,7 +206,7 @@ def run_self_forcing_tests():
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test(
|
||||
"pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -27,6 +27,11 @@ image = (
|
||||
"curl",
|
||||
"libssl-dev",
|
||||
"ffmpeg",
|
||||
"libgl1",
|
||||
"libglib2.0-0",
|
||||
"libsm6",
|
||||
"libxext6",
|
||||
"libxrender1",
|
||||
)
|
||||
.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")
|
||||
@@ -455,20 +460,31 @@ def _prepare_ssim_workspace(
|
||||
set -euo pipefail
|
||||
source $HOME/.local/bin/env
|
||||
source /opt/venv/bin/activate
|
||||
git_retry() {{
|
||||
local attempt
|
||||
for attempt in 1 2 3; do
|
||||
if "$@"; then return 0; fi
|
||||
echo "git command failed (attempt $attempt/3), retrying in 5s..."
|
||||
sleep 5
|
||||
done
|
||||
"$@"
|
||||
}}
|
||||
if [ -d {shlex.quote(repo_root)}/.git ]; then
|
||||
cd {shlex.quote(repo_root)}
|
||||
git remote set-url origin {shlex.quote(git_repo)} || true
|
||||
git fetch --prune origin
|
||||
git_retry git fetch --prune origin
|
||||
else
|
||||
git clone {shlex.quote(git_repo)} {shlex.quote(repo_root)}
|
||||
git_retry git clone {shlex.quote(git_repo)} {shlex.quote(repo_root)}
|
||||
cd {shlex.quote(repo_root)}
|
||||
fi
|
||||
{checkout_command}
|
||||
git submodule update --init --recursive
|
||||
rm -rf fastvideo/tests/ssim/reference_videos
|
||||
git_retry git submodule update --init --recursive
|
||||
cd fastvideo-kernel
|
||||
./build.sh
|
||||
cd ..
|
||||
uv pip install -e .[test]
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
export HF_HOME='/root/data/.cache'
|
||||
hf auth login --token "$HF_API_KEY"
|
||||
"""
|
||||
|
||||
BIN
Binary file not shown.
@@ -134,12 +134,11 @@ def _discover_reference_dirs_for_tier(
|
||||
|
||||
|
||||
def _has_local_reference_videos(base_dir: Path, quality_tier: str) -> bool:
|
||||
# Check for completion marker file first
|
||||
marker_path = base_dir / REFERENCE_VIDEOS_DIRNAME / quality_tier / f".download_complete_{quality_tier}"
|
||||
if not marker_path.exists():
|
||||
return False
|
||||
# Also verify at least one .mp4 exists
|
||||
for ref_dir in _discover_reference_dirs_for_tier(base_dir, quality_tier):
|
||||
tier_root = _reference_tier_root(base_dir, quality_tier)
|
||||
for ref_dir in _discover_reference_dirs(tier_root):
|
||||
for _ in _iter_video_files(ref_dir):
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
SSIM regression test for GEN3C video generation.
|
||||
|
||||
Compares newly generated GEN3C videos against device-specific reference videos
|
||||
using MS-SSIM to detect quality regressions across code changes.
|
||||
|
||||
Usage:
|
||||
# Requires 1+ GPU and reference videos.
|
||||
pytest fastvideo/tests/ssim/test_gen3c_similarity.py -v
|
||||
|
||||
Environment variables:
|
||||
GEN3C_MODEL_PATH - Diffusers-format GEN3C model path/repo id.
|
||||
Default: FastVideo/GEN3C-Cosmos-7B-Diffusers
|
||||
(local converted path also supported)
|
||||
"""
|
||||
|
||||
import os
|
||||
import glob
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
from fastvideo.worker.multiproc_executor import MultiprocExecutor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _resolve_gen3c_test_image_path() -> str:
|
||||
"""
|
||||
Resolve image path for GEN3C I2V SSIM tests.
|
||||
|
||||
Priority:
|
||||
1) GEN3C_TEST_IMAGE_PATH env var
|
||||
2) Repo asset image
|
||||
"""
|
||||
env_image = os.getenv("GEN3C_TEST_IMAGE_PATH")
|
||||
if env_image:
|
||||
return env_image
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
return str(repo_root / "assets" / "girl.png")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Device detection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
device_name = torch.cuda.get_device_name() if torch.cuda.is_available() else "cpu"
|
||||
device_reference_folder_suffix = "_reference_videos"
|
||||
|
||||
if "A40" in device_name:
|
||||
device_reference_folder = "A40" + device_reference_folder_suffix
|
||||
elif "L40S" in device_name:
|
||||
device_reference_folder = "L40S" + device_reference_folder_suffix
|
||||
else:
|
||||
device_reference_folder = None
|
||||
logger.warning(f"Unsupported device for GEN3C SSIM tests: {device_name}")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GEN3C generation parameters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
GEN3C_T2V_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": os.getenv("GEN3C_MODEL_PATH",
|
||||
"FastVideo/GEN3C-Cosmos-7B-Diffusers"),
|
||||
"height": 720,
|
||||
"width": 1280,
|
||||
"num_frames": 121,
|
||||
"num_inference_steps": 12,
|
||||
"guidance_scale": 6.0,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 1.0,
|
||||
"seed": 1024,
|
||||
"image_path": _resolve_gen3c_test_image_path(),
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"GEN3C-Cosmos-7B": GEN3C_T2V_PARAMS,
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"A camera slowly orbits around a young woman sitting at a table with a coffee mug with coffee in it in front of her, "
|
||||
"looking away naturally. Soft indoor lighting, cinematic framing, shallow depth of field, smooth camera motion.",
|
||||
]
|
||||
|
||||
BASELINE_VIDEO_NAME = "gen3c_ssim_baseline.mp4"
|
||||
CANDIDATE_VIDEO_NAME = "gen3c_ssim_candidate.mp4"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
device_reference_folder is None,
|
||||
reason=f"No reference videos for device {device_name}",
|
||||
)
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
|
||||
def test_gen3c_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"""
|
||||
Generate a GEN3C video and compare against the reference using MS-SSIM.
|
||||
"""
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
base_output_dir = os.path.join(script_dir, "generated_videos", model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = CANDIDATE_VIDEO_NAME
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
model_path = BASE_PARAMS["model_path"]
|
||||
|
||||
# Guard common misconfigurations to keep CI behavior explicit.
|
||||
if model_path.lower() == "nvidia/gen3c-cosmos-7b":
|
||||
pytest.skip(
|
||||
"nvidia/GEN3C-Cosmos-7B is the official raw checkpoint repo, not Diffusers format. "
|
||||
"Use GEN3C_MODEL_PATH=FastVideo/GEN3C-Cosmos-7B-Diffusers or a local converted path."
|
||||
)
|
||||
|
||||
local_like = model_path.startswith(("/", "./", "../"))
|
||||
if local_like and not os.path.exists(model_path):
|
||||
pytest.skip(
|
||||
f"Local GEN3C model path not found: {model_path}. "
|
||||
"Set GEN3C_MODEL_PATH to a valid local path or HF Diffusers repo id."
|
||||
)
|
||||
|
||||
if os.path.exists(model_path):
|
||||
model_index_path = os.path.join(model_path, "model_index.json")
|
||||
if not os.path.exists(model_index_path):
|
||||
pytest.skip(
|
||||
f"GEN3C_MODEL_PATH is not Diffusers-format (missing model_index.json): {model_path}"
|
||||
)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
}
|
||||
if "flow_shift" in BASE_PARAMS:
|
||||
init_kwargs["flow_shift"] = BASE_PARAMS["flow_shift"]
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"output_path": os.path.join(output_dir, output_video_name),
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"guidance_scale": BASE_PARAMS["guidance_scale"],
|
||||
"embedded_cfg_scale": BASE_PARAMS["embedded_cfg_scale"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
"image_path": BASE_PARAMS["image_path"],
|
||||
"fps": BASE_PARAMS["fps"],
|
||||
}
|
||||
|
||||
if not os.path.exists(generation_kwargs["image_path"]):
|
||||
pytest.skip(
|
||||
f"GEN3C test image not found: {generation_kwargs['image_path']}. "
|
||||
"Set GEN3C_TEST_IMAGE_PATH to a valid local image."
|
||||
)
|
||||
|
||||
# Keep local reruns deterministic: remove prior candidate outputs so
|
||||
# VideoGenerator does not auto-suffix (_1, _2, ...).
|
||||
stale_pattern = os.path.join(output_dir, "gen3c_ssim_candidate*.mp4")
|
||||
for stale_video in glob.glob(stale_pattern):
|
||||
os.remove(stale_video)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=model_path, **init_kwargs
|
||||
)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
if isinstance(generator.executor, MultiprocExecutor):
|
||||
generator.executor.shutdown()
|
||||
|
||||
assert os.path.exists(output_dir), f"Output not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(
|
||||
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
|
||||
)
|
||||
if not os.path.exists(reference_folder):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}"
|
||||
)
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, BASELINE_VIDEO_NAME)
|
||||
if not os.path.exists(reference_video_path):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video not found: {reference_video_path}"
|
||||
)
|
||||
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
|
||||
logger.info(f"Computing SSIM: {reference_video_path} vs {generated_video_path}")
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_video_path, generated_video_path, use_ms_ssim=True
|
||||
)
|
||||
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info(f"GEN3C SSIM mean: {mean_ssim}")
|
||||
|
||||
write_ssim_results(
|
||||
output_dir,
|
||||
ssim_values,
|
||||
reference_video_path,
|
||||
generated_video_path,
|
||||
num_inference_steps,
|
||||
prompt,
|
||||
)
|
||||
|
||||
# GEN3C SSIM threshold for stable L40S reference comparisons.
|
||||
min_acceptable_ssim = 0.93
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM {mean_ssim:.4f} < {min_acceptable_ssim} for {model_id} / {ATTENTION_BACKEND}"
|
||||
)
|
||||
@@ -0,0 +1,45 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.training.checkpointing_utils import ModelWrapper
|
||||
|
||||
|
||||
class DummyWrappedModule(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._lora_a = nn.Parameter(torch.tensor([1.0, 2.0]))
|
||||
self._lora_b = nn.Parameter(torch.tensor([3.0, 4.0]))
|
||||
self._frozen = nn.Parameter(torch.tensor([5.0, 6.0]), requires_grad=False)
|
||||
|
||||
def named_parameters(self, *args, **kwargs):
|
||||
# Simulate wrapped names returned by named_parameters()
|
||||
yield "layer._checkpoint_wrapped_module.lora_A", self._lora_a
|
||||
yield "layer._checkpoint_wrapped_module.lora_B", self._lora_b
|
||||
yield "layer._checkpoint_wrapped_module.frozen_weight", self._frozen
|
||||
|
||||
|
||||
def test_model_wrapper_filters_wrapped_trainable_params(monkeypatch):
|
||||
"""Regression test for wrapped parameter name mismatch during checkpoint filtering."""
|
||||
model = DummyWrappedModule()
|
||||
wrapper = ModelWrapper(model)
|
||||
|
||||
mocked_state_dict = {
|
||||
"layer.lora_A": torch.tensor([10.0, 20.0]),
|
||||
"layer.lora_B": torch.tensor([30.0, 40.0]),
|
||||
"layer.frozen_weight": torch.tensor([50.0, 60.0]),
|
||||
}
|
||||
|
||||
def mock_get_model_state_dict(_model):
|
||||
return mocked_state_dict
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.training.checkpointing_utils.get_model_state_dict",
|
||||
mock_get_model_state_dict,
|
||||
)
|
||||
|
||||
filtered_state_dict = wrapper.state_dict()
|
||||
|
||||
assert set(filtered_state_dict.keys()) == {"layer.lora_A", "layer.lora_B"}
|
||||
assert torch.equal(filtered_state_dict["layer.lora_A"], mocked_state_dict["layer.lora_A"])
|
||||
assert torch.equal(filtered_state_dict["layer.lora_B"], mocked_state_dict["layer.lora_B"])
|
||||
assert "layer.frozen_weight" not in filtered_state_dict
|
||||
+27
-11
@@ -20,12 +20,22 @@ if TYPE_CHECKING:
|
||||
TrainingConfig, )
|
||||
|
||||
|
||||
def _coerce_log_scalar(value: Any, *, where: str) -> float:
|
||||
def _coerce_log_scalar(
|
||||
value: Any,
|
||||
*,
|
||||
where: str,
|
||||
) -> float | torch.Tensor:
|
||||
"""Coerce *value* to a loggable scalar.
|
||||
|
||||
GPU tensors stay on device so we avoid a
|
||||
``cudaDeviceSynchronize`` per accumulation step.
|
||||
The caller must materialize them to float when logging.
|
||||
"""
|
||||
if isinstance(value, torch.Tensor):
|
||||
if value.numel() != 1:
|
||||
raise ValueError(f"Expected scalar tensor at {where}, "
|
||||
f"got shape={tuple(value.shape)}")
|
||||
return float(value.detach().item())
|
||||
return value.detach()
|
||||
if isinstance(value, float | int):
|
||||
return float(value)
|
||||
raise TypeError(f"Expected a scalar (float/int/Tensor) at "
|
||||
@@ -123,14 +133,16 @@ class Trainer:
|
||||
for step in progress:
|
||||
t0 = time.perf_counter()
|
||||
|
||||
loss_sums: dict[str, float] = {}
|
||||
metric_sums: dict[str, float] = {}
|
||||
# Accumulate on GPU during grad-accum; materialise
|
||||
# to CPU once per step right before logging.
|
||||
loss_sums: dict[str, float | torch.Tensor] = {}
|
||||
metric_sums: dict[str, float | torch.Tensor] = {}
|
||||
for accum_iter in range(grad_accum):
|
||||
batch = next(data_stream)
|
||||
loss_map, outputs, step_metrics = method.single_train_step(
|
||||
loss_map, outputs, step_metrics = (method.single_train_step(
|
||||
batch,
|
||||
step,
|
||||
)
|
||||
))
|
||||
|
||||
method.backward(
|
||||
loss_map,
|
||||
@@ -140,18 +152,20 @@ class Trainer:
|
||||
|
||||
for k, v in loss_map.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
loss_sums[k] = loss_sums.get(k, 0.0) + float(v.detach().item())
|
||||
prev = loss_sums.get(k, 0.0)
|
||||
loss_sums[k] = prev + v.detach()
|
||||
for k, v in step_metrics.items():
|
||||
if k in loss_sums:
|
||||
raise ValueError(f"Metric key {k!r} collides "
|
||||
"with loss key. Use a "
|
||||
"different name (e.g. prefix "
|
||||
"with 'train/').")
|
||||
metric_sums[k] = metric_sums.get(k, 0.0) + _coerce_log_scalar(
|
||||
prev = metric_sums.get(k, 0.0)
|
||||
metric_sums[k] = (prev + _coerce_log_scalar(
|
||||
v,
|
||||
where=("method.single_train_step()"
|
||||
f".metrics[{k!r}]"),
|
||||
)
|
||||
))
|
||||
|
||||
self.callbacks.on_before_optimizer_step(
|
||||
method,
|
||||
@@ -160,8 +174,10 @@ class Trainer:
|
||||
method.optimizers_schedulers_step(step)
|
||||
method.optimizers_zero_grad(step)
|
||||
|
||||
metrics = {k: v / grad_accum for k, v in loss_sums.items()}
|
||||
metrics.update({k: v / grad_accum for k, v in metric_sums.items()})
|
||||
# Single CPU sync point: materialise GPU tensors
|
||||
# to float right before logging.
|
||||
metrics = {k: float(v) / grad_accum for k, v in loss_sums.items()}
|
||||
metrics.update({k: float(v) / grad_accum for k, v in metric_sums.items()})
|
||||
metrics["step_time_sec"] = (time.perf_counter() - t0)
|
||||
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
|
||||
if self.global_rank == 0 and metrics:
|
||||
|
||||
@@ -15,11 +15,16 @@ class ModelWrapper(torch.distributed.checkpoint.stateful.Stateful):
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
state_dict = get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
# filter out non-trainable parameters
|
||||
param_requires_grad = set([k for k, v in dict(self.model.named_parameters()).items() if v.requires_grad])
|
||||
state_dict = {k: v for k, v in state_dict.items() if k in param_requires_grad}
|
||||
return state_dict # type: ignore
|
||||
state_dict = get_model_state_dict(self.model)
|
||||
|
||||
param_requires_grad = {
|
||||
k.replace("._checkpoint_wrapped_module.", ".")
|
||||
for k, v in self.model.named_parameters() if v.requires_grad
|
||||
}
|
||||
|
||||
filtered_state_dict = {k: v for k, v in state_dict.items() if k in param_requires_grad}
|
||||
|
||||
return filtered_state_dict
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
set_model_state_dict(
|
||||
|
||||
@@ -244,11 +244,21 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.latents = latents.to(
|
||||
get_local_torch_device(),
|
||||
dtype=torch.bfloat16,
|
||||
non_blocking=True,
|
||||
)
|
||||
training_batch.encoder_hidden_states = (encoder_hidden_states.to(
|
||||
get_local_torch_device(),
|
||||
dtype=torch.bfloat16,
|
||||
non_blocking=True,
|
||||
))
|
||||
training_batch.encoder_attention_mask = (encoder_attention_mask.to(
|
||||
get_local_torch_device(),
|
||||
dtype=torch.bfloat16,
|
||||
non_blocking=True,
|
||||
))
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
@@ -409,12 +419,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
|
||||
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
|
||||
# local_main_process_only=False)
|
||||
# Reduce across ranks without forcing a CPU sync
|
||||
with self.tracker.timed("timing/reduce_loss"):
|
||||
world_group = get_world_group()
|
||||
avg_loss = world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
# Accumulate on GPU; materialize to CPU only once after
|
||||
# all gradient-accumulation iterations (see train_one_step).
|
||||
training_batch.total_loss += avg_loss
|
||||
|
||||
return training_batch
|
||||
|
||||
@@ -556,7 +567,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
loss = training_batch.total_loss
|
||||
loss = float(training_batch.total_loss)
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
|
||||
+26
-5
@@ -189,10 +189,17 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
|
||||
def parse_args( # type: ignore[override]
|
||||
self, args=None, namespace=None) -> argparse.Namespace:
|
||||
namespace, unknown = self.parse_known_args(args, namespace)
|
||||
if unknown:
|
||||
self.error(f"unrecognized arguments: {' '.join(unknown)}")
|
||||
return namespace
|
||||
|
||||
def parse_known_args( # type: ignore[override]
|
||||
self, args=None, namespace=None) -> tuple[argparse.Namespace, list[str]]:
|
||||
if args is None:
|
||||
args = sys.argv[1:]
|
||||
|
||||
if '--config' in args:
|
||||
if '--config' in args and not self._should_defer_config_loading(args):
|
||||
args = self._pull_args_from_config(args)
|
||||
|
||||
# Convert underscores to dashes and vice versa in argument names
|
||||
@@ -201,10 +208,16 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
if arg.startswith('--'):
|
||||
if '=' in arg:
|
||||
key, value = arg.split('=', 1)
|
||||
key = '--' + key[len('--'):].replace('_', '-')
|
||||
normalized_key = key[len('--'):]
|
||||
if '.' not in normalized_key:
|
||||
normalized_key = normalized_key.replace('_', '-')
|
||||
key = '--' + normalized_key
|
||||
processed_args.append(f'{key}={value}')
|
||||
else:
|
||||
processed_args.append('--' + arg[len('--'):].replace('_', '-'))
|
||||
normalized_key = arg[len('--'):]
|
||||
if '.' not in normalized_key:
|
||||
normalized_key = normalized_key.replace('_', '-')
|
||||
processed_args.append('--' + normalized_key)
|
||||
elif arg.startswith('-O') and arg != '-O' and len(arg) == 2:
|
||||
# allow -O flag to be used without space, e.g. -O3
|
||||
processed_args.append('-O')
|
||||
@@ -212,7 +225,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
else:
|
||||
processed_args.append(arg)
|
||||
|
||||
namespace = super().parse_args(processed_args, namespace)
|
||||
namespace, unknown = super().parse_known_args(processed_args, namespace)
|
||||
|
||||
# Track which arguments were explicitly provided
|
||||
namespace._provided = set()
|
||||
@@ -238,7 +251,15 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
else:
|
||||
i += 1
|
||||
|
||||
return namespace # type: ignore[no-any-return]
|
||||
return namespace, unknown # type: ignore[no-any-return]
|
||||
|
||||
def _should_defer_config_loading(self, args: list[str]) -> bool:
|
||||
if getattr(self, "defer_config_loading", False):
|
||||
return True
|
||||
subcommand = next((arg for arg in args if not arg.startswith('-')), None)
|
||||
if subcommand in {"generate", "serve"}:
|
||||
return True
|
||||
return self.prog.split()[-1] in {"generate", "serve"}
|
||||
|
||||
def _pull_args_from_config(self, args: list[str]) -> list[str]:
|
||||
"""Method to pull arguments specified in the config file
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable
|
||||
from queue import Queue
|
||||
from typing import Any, TypeVar, cast
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -14,8 +15,14 @@ _R = TypeVar("_R")
|
||||
|
||||
class Executor(ABC):
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs):
|
||||
def __init__(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
*,
|
||||
log_queue=None,
|
||||
):
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self._log_queue = log_queue
|
||||
|
||||
self._init_executor()
|
||||
|
||||
@@ -97,6 +104,16 @@ class Executor(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def set_log_queue(self, log_queue: Queue | None) -> None:
|
||||
"""Forward worker logs to the given queue. Call before generate_video."""
|
||||
self.collective_rpc("set_log_queue", kwargs={"log_queue": log_queue})
|
||||
|
||||
@abstractmethod
|
||||
def clear_log_queue(self) -> None:
|
||||
"""Stop forwarding worker logs to the queue. Call after generate_video."""
|
||||
self.collective_rpc("clear_log_queue")
|
||||
|
||||
@abstractmethod
|
||||
def shutdown(self) -> None:
|
||||
"""
|
||||
|
||||
@@ -7,6 +7,8 @@ import contextlib
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
import faulthandler
|
||||
import logging
|
||||
import logging.handlers
|
||||
import multiprocessing as mp
|
||||
from multiprocessing.connection import Connection
|
||||
from multiprocessing.queues import Queue
|
||||
@@ -34,6 +36,11 @@ from fastvideo.worker.worker_base import WorkerWrapperBase
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _make_queue_log_handler(log_queue: Queue) -> logging.Handler:
|
||||
"""Create a QueueHandler that forwards fastvideo logs to a multiprocessing queue."""
|
||||
return logging.handlers.QueueHandler(log_queue)
|
||||
|
||||
|
||||
class StreamingTaskType(str, Enum):
|
||||
"""
|
||||
Enumeration for different streaming task types.
|
||||
@@ -97,6 +104,7 @@ class MultiprocExecutor(Executor):
|
||||
distributed_init_method=distributed_init_method,
|
||||
streaming_input_queue=self._streaming_input_queue,
|
||||
streaming_output_queue=self._streaming_output_queue,
|
||||
log_queue=self._log_queue,
|
||||
))
|
||||
|
||||
# Workers must be created before wait_for_ready to avoid
|
||||
@@ -246,6 +254,14 @@ class MultiprocExecutor(Executor):
|
||||
if response["status"] != "lora_adapter_merged":
|
||||
raise RuntimeError(f"Worker {i} failed to merge LoRA weights")
|
||||
|
||||
def set_log_queue(self, log_queue: Queue | None) -> None:
|
||||
"""Forward worker logs to the given queue. Call before generate_video."""
|
||||
self.collective_rpc("set_log_queue", kwargs={"log_queue": log_queue})
|
||||
|
||||
def clear_log_queue(self) -> None:
|
||||
"""Stop forwarding worker logs to the queue. Call after generate_video."""
|
||||
self.collective_rpc("clear_log_queue")
|
||||
|
||||
def collective_rpc(self,
|
||||
method: str | Callable,
|
||||
timeout: float | None = None,
|
||||
@@ -298,6 +314,12 @@ class MultiprocExecutor(Executor):
|
||||
return # Prevent multiple shutdown calls
|
||||
|
||||
logger.info("Shutting down MultiprocExecutor...")
|
||||
|
||||
# Check if workers were initialized (they might not be if initialization failed)
|
||||
if not hasattr(self, 'workers') or not self.workers:
|
||||
logger.info("No workers to shut down.")
|
||||
return
|
||||
|
||||
self.shutting_down = True
|
||||
|
||||
# First try gentle termination
|
||||
@@ -427,11 +449,14 @@ class WorkerMultiprocProc:
|
||||
pipe: Connection,
|
||||
streaming_input_queue: Queue | None = None,
|
||||
streaming_output_queue: Queue | None = None,
|
||||
_initial_log_handler: logging.Handler | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
self.rank = rank
|
||||
self.pipe = pipe
|
||||
self.streaming_input_queue = streaming_input_queue
|
||||
self.streaming_output_queue = streaming_output_queue
|
||||
self._initial_log_handler = _initial_log_handler
|
||||
wrapper = WorkerWrapperBase(fastvideo_args=fastvideo_args, rpc_rank=rank)
|
||||
|
||||
all_kwargs: list[dict] = [{} for _ in range(fastvideo_args.num_gpus)]
|
||||
@@ -458,6 +483,7 @@ class WorkerMultiprocProc:
|
||||
distributed_init_method: str,
|
||||
streaming_input_queue: Queue | None = None,
|
||||
streaming_output_queue: Queue | None = None,
|
||||
log_queue: Queue | None = None,
|
||||
) -> UnreadyWorkerProcHandle:
|
||||
context = get_mp_context()
|
||||
executor_pipe, worker_pipe = context.Pipe(duplex=True)
|
||||
@@ -472,6 +498,7 @@ class WorkerMultiprocProc:
|
||||
"ready_pipe": writer,
|
||||
"streaming_input_queue": streaming_input_queue,
|
||||
"streaming_output_queue": streaming_output_queue,
|
||||
"log_queue": log_queue,
|
||||
}
|
||||
# Run EngineCore busy loop in background process.
|
||||
proc = context.Process(target=WorkerMultiprocProc.worker_main,
|
||||
@@ -488,6 +515,13 @@ class WorkerMultiprocProc:
|
||||
""" Worker initialization and execution loops.
|
||||
This runs a background process """
|
||||
|
||||
log_queue = kwargs.pop("log_queue", None)
|
||||
# Add log handler before model loading so we capture fsdp_load, cuda, etc.
|
||||
if log_queue is not None:
|
||||
_handler = _make_queue_log_handler(log_queue)
|
||||
logging.getLogger("fastvideo").addHandler(_handler)
|
||||
kwargs["_initial_log_handler"] = _handler
|
||||
|
||||
# Signal handler used for graceful termination.
|
||||
# SystemExit exception is only raised once to allow this and worker
|
||||
# processes to terminate without error
|
||||
@@ -523,9 +557,21 @@ class WorkerMultiprocProc:
|
||||
|
||||
worker.worker_busy_loop()
|
||||
|
||||
except Exception:
|
||||
except Exception as exc:
|
||||
if ready_pipe is not None:
|
||||
logger.exception("WorkerMultiprocProc failed to start.")
|
||||
# Send error status to parent before closing pipe
|
||||
try:
|
||||
traceback_str = get_exception_traceback()
|
||||
ready_pipe.send({
|
||||
"status": "ERROR",
|
||||
"error": str(exc),
|
||||
"traceback": traceback_str,
|
||||
"rank": rank,
|
||||
})
|
||||
except Exception:
|
||||
# If sending fails, at least log it
|
||||
pass
|
||||
else:
|
||||
logger.exception("WorkerMultiprocProc failed.")
|
||||
|
||||
@@ -535,7 +581,8 @@ class WorkerMultiprocProc:
|
||||
shutdown_requested = True
|
||||
traceback = get_exception_traceback()
|
||||
logger.error("Worker %d hit an exception: %s", rank, traceback)
|
||||
parent_process.send_signal(signal.SIGQUIT)
|
||||
if parent_process:
|
||||
parent_process.send_signal(signal.SIGQUIT)
|
||||
|
||||
finally:
|
||||
if ready_pipe is not None:
|
||||
@@ -552,7 +599,8 @@ class WorkerMultiprocProc:
|
||||
"See stack trace for root cause.")
|
||||
|
||||
pipes = {handle.ready_pipe: handle for handle in unready_proc_handles}
|
||||
ready_proc_handles: list[WorkerProcHandle | None] = ([None] * len(unready_proc_handles))
|
||||
ready_proc_handles: list[WorkerProcHandle | None] = [None] * len(unready_proc_handles)
|
||||
worker_errors: list[str] = []
|
||||
while pipes:
|
||||
ready = mp.connection.wait(pipes.keys())
|
||||
for pipe in ready:
|
||||
@@ -561,13 +609,28 @@ class WorkerMultiprocProc:
|
||||
# Wait until the WorkerProc is ready.
|
||||
unready_proc_handle = pipes.pop(pipe)
|
||||
response: dict[str, Any] = pipe.recv()
|
||||
if response["status"] != "READY":
|
||||
raise e
|
||||
if response["status"] == "ERROR":
|
||||
# Worker sent error details
|
||||
error_msg = response.get("error", "Unknown error")
|
||||
traceback_str = response.get("traceback", "")
|
||||
rank = response.get("rank", "unknown")
|
||||
error_info = f"Worker {rank} error: {error_msg}"
|
||||
if traceback_str:
|
||||
error_info += f"\n{traceback_str}"
|
||||
worker_errors.append(error_info)
|
||||
# Log a concise error message (full traceback will be in the exception)
|
||||
logger.error("Worker %s initialization failed: %s", rank, error_msg)
|
||||
# Continue to check other workers, but we'll fail at the end
|
||||
elif response["status"] != "READY":
|
||||
worker_errors.append(f"Worker returned unexpected status: {response.get('status', 'unknown')}")
|
||||
|
||||
ready_proc_handles[unready_proc_handle.rank] = (
|
||||
WorkerProcHandle.from_unready_handle(unready_proc_handle))
|
||||
if response["status"] == "READY":
|
||||
ready_proc_handles[unready_proc_handle.rank] = (
|
||||
WorkerProcHandle.from_unready_handle(unready_proc_handle))
|
||||
|
||||
except EOFError:
|
||||
# Pipe closed without sending status - worker crashed
|
||||
worker_errors.append("Worker process crashed (pipe closed unexpectedly)")
|
||||
e.__suppress_context__ = True
|
||||
raise e from None
|
||||
|
||||
@@ -575,6 +638,12 @@ class WorkerMultiprocProc:
|
||||
# Close connection.
|
||||
pipe.close()
|
||||
|
||||
# If any workers failed, raise exception with details
|
||||
if worker_errors:
|
||||
error_msg = "WorkerMultiprocProc initialization failed due to exceptions in background processes:\n"
|
||||
error_msg += "\n".join(f" - {err}" for err in worker_errors)
|
||||
raise Exception(error_msg) from None
|
||||
|
||||
logger.info("%d workers ready", len(ready_proc_handles))
|
||||
return cast(list[WorkerProcHandle], ready_proc_handles)
|
||||
|
||||
@@ -597,6 +666,14 @@ class WorkerMultiprocProc:
|
||||
with contextlib.suppress(Exception):
|
||||
self.pipe.send(response)
|
||||
break
|
||||
if method == "set_log_queue":
|
||||
self._set_log_queue(kwargs.get("log_queue"))
|
||||
self.pipe.send({"status": "ok"})
|
||||
continue
|
||||
if method == "clear_log_queue":
|
||||
self._clear_log_queue()
|
||||
self.pipe.send({"status": "ok"})
|
||||
continue
|
||||
if method == "start_streaming_queue_loop":
|
||||
self.pipe.send({"status": "streaming_queue_loop_started"})
|
||||
self.streaming_queue_loop()
|
||||
@@ -665,6 +742,26 @@ class WorkerMultiprocProc:
|
||||
logger.error("Worker %d queue loop error: %s", self.rank, e)
|
||||
self.streaming_output_queue.put(StreamingResult(task_type=StreamingTaskType.STEP, error=e))
|
||||
|
||||
_log_queue_handler: logging.Handler | None = None
|
||||
|
||||
def _set_log_queue(self, log_queue: Queue | None) -> None:
|
||||
"""Add a handler that forwards fastvideo logs to the given queue."""
|
||||
self._clear_log_queue()
|
||||
if log_queue is None:
|
||||
return
|
||||
# Remove initial handler if present (from worker_main) to avoid duplicates
|
||||
if self._initial_log_handler is not None:
|
||||
logging.getLogger("fastvideo").removeHandler(self._initial_log_handler)
|
||||
self._initial_log_handler = None
|
||||
self._log_queue_handler = _make_queue_log_handler(log_queue)
|
||||
logging.getLogger("fastvideo").addHandler(self._log_queue_handler)
|
||||
|
||||
def _clear_log_queue(self) -> None:
|
||||
"""Remove the log queue handler."""
|
||||
if self._log_queue_handler is not None:
|
||||
logging.getLogger("fastvideo").removeHandler(self._log_queue_handler)
|
||||
self._log_queue_handler = None
|
||||
|
||||
@staticmethod
|
||||
def setup_proc_title_and_log_prefix() -> None:
|
||||
dp_size = get_dp_group().world_size
|
||||
|
||||
+5
-1
@@ -80,6 +80,10 @@ dependencies = [
|
||||
"torchcodec",
|
||||
"ray>=2.49.1",
|
||||
"ftfy==6.3.1",
|
||||
|
||||
# Job Runner
|
||||
"fastapi==0.129.0",
|
||||
"uvicorn==0.41.0"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -146,7 +150,7 @@ check_untyped_defs = true
|
||||
follow_imports = "silent"
|
||||
|
||||
[tool.codespell]
|
||||
skip ="./data,./wandb"
|
||||
skip = "./data,./wandb,ui/package-lock.json"
|
||||
|
||||
[tool.ruff]
|
||||
# Allow lines to be as long as 120.
|
||||
|
||||
@@ -0,0 +1,724 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Convert GEN3C checkpoint (nvidia/GEN3C-Cosmos-7B) to FastVideo (diffusers) format.
|
||||
|
||||
Usage:
|
||||
# Convert from local checkpoint (memory-efficient mode for limited RAM)
|
||||
python convert_gen3c_to_fastvideo.py --source ./official_weights/GEN3C-Cosmos-7B/model.pt --output ./gen3c_fastvideo
|
||||
|
||||
# Download and convert from HuggingFace
|
||||
python convert_gen3c_to_fastvideo.py --download nvidia/GEN3C-Cosmos-7B --output ./gen3c_fastvideo
|
||||
|
||||
# Analyze checkpoint structure only (low memory)
|
||||
python convert_gen3c_to_fastvideo.py --source ./model.pt --analyze
|
||||
|
||||
# Convert with fp16 to reduce output size (and memory during save)
|
||||
python convert_gen3c_to_fastvideo.py --source ./model.pt --output ./gen3c_fastvideo --dtype fp16
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Iterator
|
||||
|
||||
import torch
|
||||
from safetensors.torch import save_file
|
||||
|
||||
try:
|
||||
from huggingface_hub import hf_hub_download, snapshot_download
|
||||
except ImportError:
|
||||
hf_hub_download = None
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
# Parameter name mapping from official GEN3C checkpoint to FastVideo format
|
||||
# Based on fastvideo/configs/models/dits/gen3c.py
|
||||
# The actual GEN3C checkpoint uses AdaLN-LoRA with decomposed projections (index 0 = base, 1 = LoRA)
|
||||
PARAM_NAMES_MAPPING: dict[str, str] = {
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding
|
||||
r"^net\.t_embedder\.0\.(.*)$": r"time_embed.time_proj.\1",
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$": r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$": r"time_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Augment sigma embedding (GEN3C-specific)
|
||||
r"^net\.augment_sigma_embedder\.0\.(.*)$": r"augment_sigma_embed.time_proj.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_1\.(.*)$": r"augment_sigma_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_2\.(.*)$": r"augment_sigma_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Affine embedding norm
|
||||
r"^net\.affline_norm\.(.*)$": r"affine_norm.\1",
|
||||
|
||||
# Extra positional embeddings (learnable per-axis)
|
||||
r"^net\.extra_pos_embedder\.pos_emb_t$": r"learnable_pos_embed.pos_emb_t",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_h$": r"learnable_pos_embed.pos_emb_h",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_w$": r"learnable_pos_embed.pos_emb_w",
|
||||
|
||||
# Transformer blocks: net.blocks.blockN -> transformer_blocks.N
|
||||
# GEN3C uses attn.to_q.0/1 pattern (0=base weight, 1=LoRA weight)
|
||||
|
||||
# Self-attention (block index 0)
|
||||
# Q projection: to_q.0 is linear, to_q.1 is QK norm (RMSNorm applied per-head)
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.1\.(.*)$": r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
# K projection: to_k.0 is linear, to_k.1 is QK norm
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.1\.(.*)$": r"transformer_blocks.\1.attn1.norm_k.\2",
|
||||
# V projection
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
# Output projection
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_out\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
# AdaLN modulation for self-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.adaLN_modulation\.(.*)$": r"transformer_blocks.\1.adaln_modulation_self_attn.\2",
|
||||
|
||||
# Cross-attention (block index 1)
|
||||
# Q projection: to_q.0 is linear, to_q.1 is QK norm
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.1\.(.*)$": r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
# K projection: to_k.0 is linear, to_k.1 is QK norm
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.1\.(.*)$": r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
# V projection
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
# Output projection
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_out\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
# AdaLN modulation for cross-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.adaLN_modulation\.(.*)$": r"transformer_blocks.\1.adaln_modulation_cross_attn.\2",
|
||||
|
||||
# MLP (block index 2) - simpler naming: layer1, layer2 directly
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.adaLN_modulation\.(.*)$": r"transformer_blocks.\1.adaln_modulation_mlp.\2",
|
||||
|
||||
# Final layer
|
||||
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
|
||||
r"^net\.final_layer\.adaLN_modulation\.(.*)$": r"final_layer.adaln_modulation.\1",
|
||||
}
|
||||
|
||||
# Keys to skip (dynamically computed or training metadata)
|
||||
SKIP_PATTERNS = [
|
||||
"net.pos_embedder.", # RoPE computed dynamically
|
||||
"net.accum_", # Training accumulation metadata
|
||||
"logvar.", # Training-only logvar module (not used for inference)
|
||||
]
|
||||
|
||||
|
||||
def apply_mapping(key: str) -> str | None:
|
||||
"""Apply parameter name mapping to convert from official to FastVideo format."""
|
||||
# Check if key should be skipped
|
||||
for pattern in SKIP_PATTERNS:
|
||||
if key.startswith(pattern):
|
||||
return None
|
||||
|
||||
# Apply mapping patterns
|
||||
for pattern, replacement in PARAM_NAMES_MAPPING.items():
|
||||
if re.match(pattern, key):
|
||||
return re.sub(pattern, replacement, key)
|
||||
|
||||
# If no mapping found, return original key (will be reported)
|
||||
return key
|
||||
|
||||
|
||||
def load_checkpoint(path: Path, mmap: bool = True) -> dict[str, torch.Tensor]:
|
||||
"""Load checkpoint from .pt file.
|
||||
|
||||
Args:
|
||||
path: Path to checkpoint file
|
||||
mmap: Use memory-mapped loading (PyTorch 2.1+) for reduced RAM usage
|
||||
"""
|
||||
# Try memory-mapped loading first (PyTorch 2.1+)
|
||||
load_kwargs: dict = {"map_location": "cpu", "weights_only": False}
|
||||
|
||||
if mmap:
|
||||
# Check if PyTorch version supports mmap (2.1+)
|
||||
try:
|
||||
version_str = torch.__version__.split("+")[0] # Remove +cu118 suffix etc.
|
||||
major, minor = version_str.split(".")[:2]
|
||||
torch_version = (int(major), int(minor))
|
||||
if torch_version >= (2, 1):
|
||||
load_kwargs["mmap"] = True
|
||||
print(" Using memory-mapped loading for reduced RAM usage")
|
||||
else:
|
||||
print(f" Note: PyTorch {torch.__version__} doesn't support mmap, using standard loading")
|
||||
except (ValueError, IndexError):
|
||||
print(f" Warning: Could not parse PyTorch version {torch.__version__}, skipping mmap")
|
||||
|
||||
checkpoint = torch.load(path, **load_kwargs)
|
||||
|
||||
# Handle different checkpoint formats
|
||||
if isinstance(checkpoint, dict):
|
||||
for key in ("state_dict", "model_state_dict", "model", "ema"):
|
||||
if key in checkpoint:
|
||||
return checkpoint[key]
|
||||
return checkpoint
|
||||
|
||||
return checkpoint
|
||||
|
||||
|
||||
def iterate_checkpoint_keys(path: Path) -> Iterator[str]:
|
||||
"""Iterate over checkpoint keys without loading all tensors.
|
||||
|
||||
This is useful for analysis when memory is limited.
|
||||
"""
|
||||
# Load with weights_only=True to just get the structure
|
||||
# Unfortunately PyTorch doesn't have a great way to do this
|
||||
# So we load the full checkpoint but only keep keys
|
||||
checkpoint = torch.load(path, map_location="meta", weights_only=False)
|
||||
|
||||
if isinstance(checkpoint, dict):
|
||||
for key in ("state_dict", "model_state_dict", "model", "ema"):
|
||||
if key in checkpoint:
|
||||
checkpoint = checkpoint[key]
|
||||
break
|
||||
|
||||
return checkpoint.keys()
|
||||
|
||||
|
||||
def analyze_checkpoint(state_dict: dict[str, torch.Tensor]) -> None:
|
||||
"""Print checkpoint structure analysis."""
|
||||
print("\n" + "=" * 80)
|
||||
print("CHECKPOINT ANALYSIS")
|
||||
print("=" * 80)
|
||||
|
||||
# Count parameters
|
||||
total_params = sum(p.numel() for p in state_dict.values())
|
||||
print(f"\nTotal parameters: {total_params:,} ({total_params / 1e9:.2f}B)")
|
||||
print(f"Total keys: {len(state_dict)}")
|
||||
|
||||
# Analyze key prefixes
|
||||
prefixes: dict[str, int] = {}
|
||||
for key in state_dict:
|
||||
parts = key.split(".")
|
||||
prefix = ".".join(parts[:2]) if len(parts) > 1 else parts[0]
|
||||
prefixes[prefix] = prefixes.get(prefix, 0) + 1
|
||||
|
||||
print("\nKey prefixes:")
|
||||
for prefix, count in sorted(prefixes.items()):
|
||||
print(f" {prefix}: {count}")
|
||||
|
||||
# Print first 100 keys with shapes
|
||||
print("\nFirst 100 keys with shapes:")
|
||||
for i, (key, value) in enumerate(state_dict.items()):
|
||||
if i >= 100:
|
||||
print(f" ... and {len(state_dict) - 100} more keys")
|
||||
break
|
||||
print(f" {key}: {list(value.shape)}")
|
||||
|
||||
# Identify GEN3C-specific layers
|
||||
print("\nGEN3C-specific layers:")
|
||||
gen3c_patterns = [
|
||||
"augment_sigma",
|
||||
"x_embedder", # Has more input channels than standard Cosmos
|
||||
"extra_pos_embedder",
|
||||
]
|
||||
for key in state_dict:
|
||||
for pattern in gen3c_patterns:
|
||||
if pattern in key:
|
||||
print(f" {key}: {list(state_dict[key].shape)}")
|
||||
break
|
||||
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
|
||||
def convert_weights(
|
||||
state_dict: dict[str, torch.Tensor],
|
||||
verbose: bool = False,
|
||||
dtype: torch.dtype | None = None,
|
||||
memory_efficient: bool = True,
|
||||
) -> tuple[OrderedDict[str, torch.Tensor], list[str], list[str]]:
|
||||
"""Convert weights from official format to FastVideo format.
|
||||
|
||||
Args:
|
||||
state_dict: Source state dict
|
||||
verbose: Print detailed conversion info
|
||||
dtype: Convert tensors to this dtype (e.g., torch.float16)
|
||||
memory_efficient: Delete source tensors after processing to free memory
|
||||
|
||||
Returns:
|
||||
converted: Converted state dict
|
||||
unmapped: List of keys that weren't mapped (kept as-is)
|
||||
skipped: List of keys that were skipped
|
||||
"""
|
||||
converted = OrderedDict()
|
||||
unmapped = []
|
||||
skipped = []
|
||||
|
||||
# Get all keys first (so we can delete as we go)
|
||||
keys = list(state_dict.keys())
|
||||
total = len(keys)
|
||||
|
||||
for i, key in enumerate(keys):
|
||||
value = state_dict[key]
|
||||
new_key = apply_mapping(key)
|
||||
|
||||
if new_key is None:
|
||||
skipped.append(key)
|
||||
if verbose:
|
||||
print(f" Skipped: {key}")
|
||||
elif new_key == key:
|
||||
# No mapping found, but not in skip list
|
||||
unmapped.append(key)
|
||||
if dtype is not None and value.is_floating_point():
|
||||
converted[key] = value.to(dtype).contiguous()
|
||||
else:
|
||||
converted[key] = value.contiguous() if memory_efficient else value
|
||||
if verbose:
|
||||
print(f" Unmapped: {key}")
|
||||
else:
|
||||
if dtype is not None and value.is_floating_point():
|
||||
converted[new_key] = value.to(dtype).contiguous()
|
||||
else:
|
||||
converted[new_key] = value.contiguous() if memory_efficient else value
|
||||
if verbose:
|
||||
print(f" {key} -> {new_key}")
|
||||
|
||||
# Free memory by deleting processed tensor from source
|
||||
if memory_efficient:
|
||||
del state_dict[key]
|
||||
if i % 100 == 0:
|
||||
gc.collect()
|
||||
|
||||
# Progress indicator
|
||||
if (i + 1) % 200 == 0 or i == total - 1:
|
||||
print(f" Processed {i + 1}/{total} tensors...")
|
||||
|
||||
# Final garbage collection
|
||||
if memory_efficient:
|
||||
gc.collect()
|
||||
|
||||
return converted, unmapped, skipped
|
||||
|
||||
|
||||
def write_component(
|
||||
output_dir: Path,
|
||||
name: str,
|
||||
weights: OrderedDict[str, torch.Tensor],
|
||||
config: dict | None = None,
|
||||
) -> None:
|
||||
"""Write component weights and config to output directory."""
|
||||
component_dir = output_dir / name
|
||||
component_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Save weights
|
||||
output_file = component_dir / "model.safetensors"
|
||||
save_file(weights, str(output_file))
|
||||
print(f"Saved {name} weights to {output_file}")
|
||||
print(f" {len(weights)} tensors, {sum(t.numel() for t in weights.values()):,} parameters")
|
||||
|
||||
# Save config
|
||||
if config is not None:
|
||||
config_path = component_dir / "config.json"
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(config, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"Saved {name} config to {config_path}")
|
||||
|
||||
|
||||
def build_transformer_config() -> dict:
|
||||
"""Build transformer config for Gen3CTransformer3DModel.
|
||||
|
||||
Architecture based on checkpoint analysis:
|
||||
- hidden_size = 4096 (32 heads * 128 head_dim)
|
||||
- num_layers = 28
|
||||
- in_channels = 82 (16 VAE + 1 mask + 64 buffer + 1 padding)
|
||||
- MLP ratio = 4.0
|
||||
"""
|
||||
return {
|
||||
"_class_name": "Gen3CTransformer3DModel",
|
||||
"in_channels": 16, # Base VAE channels (full input computed at runtime)
|
||||
"out_channels": 16,
|
||||
"num_attention_heads": 32, # 4096 / 128 = 32
|
||||
"attention_head_dim": 128,
|
||||
"num_layers": 28,
|
||||
"mlp_ratio": 4.0,
|
||||
"text_embed_dim": 1024, # T5 embedding dim
|
||||
"adaln_lora_dim": 256,
|
||||
"use_adaln_lora": True,
|
||||
"add_augment_sigma_embedding": False, # Not present in this checkpoint
|
||||
"frame_buffer_max": 2, # 2 buffers for 3D cache
|
||||
"max_size": [128, 240, 240], # Max T, H, W for positional embeddings
|
||||
"patch_size": [1, 2, 2],
|
||||
"rope_scale": [2.0, 1.0, 1.0],
|
||||
"extra_pos_embed_type": "learnable",
|
||||
"concat_padding_mask": True,
|
||||
"affine_emb_norm": True,
|
||||
"qk_norm": "rms_norm",
|
||||
"eps": 1e-6,
|
||||
}
|
||||
|
||||
|
||||
def build_model_index() -> dict:
|
||||
"""Build model_index.json for the converted model."""
|
||||
return {
|
||||
"_class_name": "Gen3CPipeline",
|
||||
"_diffusers_version": "0.33.0.dev0",
|
||||
"transformer": ["diffusers", "Gen3CTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLGen3CTokenizer"],
|
||||
"text_encoder": ["transformers", "T5EncoderModel"],
|
||||
"tokenizer": ["transformers", "T5Tokenizer"],
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
}
|
||||
|
||||
|
||||
def download_checkpoint(
|
||||
repo_id: str,
|
||||
filename: str = "model.pt",
|
||||
token: str | None = None,
|
||||
cache_dir: Path | None = None,
|
||||
) -> Path:
|
||||
"""Download checkpoint from HuggingFace Hub."""
|
||||
if hf_hub_download is None:
|
||||
raise RuntimeError("huggingface_hub is required for --download. Install with: pip install huggingface_hub")
|
||||
|
||||
print(f"Downloading {filename} from {repo_id}...")
|
||||
path = hf_hub_download(
|
||||
repo_id=repo_id,
|
||||
filename=filename,
|
||||
token=token,
|
||||
cache_dir=str(cache_dir) if cache_dir else None,
|
||||
)
|
||||
print(f"Downloaded to {path}")
|
||||
return Path(path)
|
||||
|
||||
|
||||
def resolve_model_dir(
|
||||
model_name_or_path: str,
|
||||
cache_dir: Path | None = None,
|
||||
) -> Path:
|
||||
"""Resolve a local model directory from path or HuggingFace repo id."""
|
||||
model_path = Path(model_name_or_path)
|
||||
if model_path.exists():
|
||||
return model_path
|
||||
|
||||
if snapshot_download is None:
|
||||
raise RuntimeError(
|
||||
"huggingface_hub is required to download component source. "
|
||||
"Install with: pip install huggingface_hub")
|
||||
|
||||
print(f"Downloading component source repo: {model_name_or_path}")
|
||||
downloaded = snapshot_download(
|
||||
repo_id=model_name_or_path,
|
||||
allow_patterns=[
|
||||
"model_index.json",
|
||||
"vae/*",
|
||||
"text_encoder/*",
|
||||
"tokenizer/*",
|
||||
"scheduler/*",
|
||||
],
|
||||
local_dir=str(cache_dir) if cache_dir else None,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
print(f"Component source downloaded to {downloaded}")
|
||||
return Path(downloaded)
|
||||
|
||||
|
||||
def add_inference_components(
|
||||
source_dir: Path,
|
||||
output_dir: Path,
|
||||
link_components: bool = False,
|
||||
) -> None:
|
||||
"""Copy or symlink VAE/text encoder/tokenizer/scheduler into output dir."""
|
||||
required_components = ("vae", "text_encoder", "tokenizer", "scheduler")
|
||||
missing: list[str] = []
|
||||
|
||||
for component in required_components:
|
||||
src = source_dir / component
|
||||
dst = output_dir / component
|
||||
|
||||
if not src.exists():
|
||||
missing.append(component)
|
||||
continue
|
||||
|
||||
if dst.exists():
|
||||
print(f" Skipping {component}: already exists at {dst}")
|
||||
continue
|
||||
|
||||
if link_components:
|
||||
dst.symlink_to(src.resolve(), target_is_directory=True)
|
||||
print(f" Linked {component}: {dst} -> {src.resolve()}")
|
||||
else:
|
||||
shutil.copytree(src, dst, dirs_exist_ok=False)
|
||||
print(f" Copied {component}: {src} -> {dst}")
|
||||
|
||||
if missing:
|
||||
raise FileNotFoundError(
|
||||
f"Missing required components in source repo {source_dir}: {missing}"
|
||||
)
|
||||
|
||||
|
||||
def patch_gen3c_vae_config(output_dir: Path) -> None:
|
||||
"""Tag copied VAE config with the Gen3C tokenizer-backed class name."""
|
||||
vae_cfg_path = output_dir / "vae" / "config.json"
|
||||
if not vae_cfg_path.exists():
|
||||
return
|
||||
with vae_cfg_path.open("r", encoding="utf-8") as f:
|
||||
cfg = json.load(f)
|
||||
cfg["_class_name"] = "AutoencoderKLGen3CTokenizer"
|
||||
with vae_cfg_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(cfg, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f" Patched VAE config class to AutoencoderKLGen3CTokenizer: {vae_cfg_path}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Convert GEN3C checkpoint to FastVideo format",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
epilog="""
|
||||
Examples:
|
||||
# Convert from local checkpoint (recommended for limited RAM)
|
||||
python convert_gen3c_to_fastvideo.py --source ./official_weights/GEN3C-Cosmos-7B/model.pt --output ./gen3c_fastvideo
|
||||
|
||||
# Download and convert from HuggingFace
|
||||
python convert_gen3c_to_fastvideo.py --download nvidia/GEN3C-Cosmos-7B --output ./gen3c_fastvideo
|
||||
|
||||
# Analyze checkpoint structure only
|
||||
python convert_gen3c_to_fastvideo.py --source ./model.pt --analyze
|
||||
|
||||
# Convert to fp16 for smaller output (and lower memory during save)
|
||||
python convert_gen3c_to_fastvideo.py --source ./model.pt --output ./gen3c_fastvideo --dtype fp16
|
||||
""",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--source",
|
||||
type=str,
|
||||
help="Path to input .pt checkpoint file",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
help="Output directory for converted weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--download",
|
||||
type=str,
|
||||
help="HuggingFace repo ID to download checkpoint from (e.g., nvidia/GEN3C-Cosmos-7B)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--filename",
|
||||
type=str,
|
||||
default="model.pt",
|
||||
help="Filename to download from HuggingFace (default: model.pt)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--token",
|
||||
type=str,
|
||||
default=os.getenv("HF_TOKEN"),
|
||||
help="HuggingFace token (or set HF_TOKEN env var)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--analyze",
|
||||
action="store_true",
|
||||
help="Only analyze checkpoint structure, don't convert",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verbose",
|
||||
action="store_true",
|
||||
help="Print detailed conversion info",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--force",
|
||||
action="store_true",
|
||||
help="Overwrite output directory if it exists",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
default="bf16",
|
||||
help="Output dtype for weights (default: bf16, use fp16 for smaller output)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-mmap",
|
||||
action="store_true",
|
||||
help="Disable memory-mapped loading (not recommended, uses more RAM)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--components-source",
|
||||
type=str,
|
||||
default=None,
|
||||
help=(
|
||||
"Optional local path or HF repo id containing diffusers components "
|
||||
"(vae/text_encoder/tokenizer/scheduler) to copy into output. "
|
||||
"Example: nvidia/Cosmos-Predict2-2B-Video2World"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--components-cache-dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Optional cache/local directory used when downloading --components-source.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--link-components",
|
||||
action="store_true",
|
||||
help="Create symlinks for components instead of copying (for local sources).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--components-only",
|
||||
action="store_true",
|
||||
help=(
|
||||
"Skip transformer conversion and only add "
|
||||
"vae/text_encoder/tokenizer/scheduler into --output."
|
||||
),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Validate arguments
|
||||
if args.components_only:
|
||||
if args.output is None:
|
||||
raise ValueError("--output is required with --components-only")
|
||||
if args.components_source is None:
|
||||
raise ValueError(
|
||||
"--components-source is required with --components-only")
|
||||
else:
|
||||
if args.download and args.source:
|
||||
raise ValueError("Use either --download or --source, not both")
|
||||
if not args.download and not args.source:
|
||||
raise ValueError("Either --download or --source is required")
|
||||
if not args.analyze and not args.output:
|
||||
raise ValueError("--output is required when not using --analyze")
|
||||
|
||||
if args.components_only:
|
||||
output_dir = Path(args.output)
|
||||
if not output_dir.exists():
|
||||
raise FileNotFoundError(
|
||||
f"--components-only expected existing output directory: {output_dir}"
|
||||
)
|
||||
component_source_dir = resolve_model_dir(
|
||||
args.components_source,
|
||||
cache_dir=Path(args.components_cache_dir)
|
||||
if args.components_cache_dir else None,
|
||||
)
|
||||
print("Adding inference components only...")
|
||||
add_inference_components(
|
||||
source_dir=component_source_dir,
|
||||
output_dir=output_dir,
|
||||
link_components=args.link_components,
|
||||
)
|
||||
patch_gen3c_vae_config(output_dir)
|
||||
print(f"Done. Components added to {output_dir}")
|
||||
return
|
||||
|
||||
# Parse dtype
|
||||
dtype_map = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
target_dtype = dtype_map[args.dtype]
|
||||
|
||||
# Get checkpoint path
|
||||
if args.download:
|
||||
checkpoint_path = download_checkpoint(
|
||||
repo_id=args.download,
|
||||
filename=args.filename,
|
||||
token=args.token,
|
||||
)
|
||||
else:
|
||||
checkpoint_path = Path(args.source)
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
|
||||
# Load checkpoint with memory-mapped loading
|
||||
print(f"Loading checkpoint from {checkpoint_path}...")
|
||||
use_mmap = not args.no_mmap
|
||||
state_dict = load_checkpoint(checkpoint_path, mmap=use_mmap)
|
||||
|
||||
# Analyze if requested
|
||||
if args.analyze:
|
||||
analyze_checkpoint(state_dict)
|
||||
return
|
||||
|
||||
# Analyze before conversion (quick summary only to save memory)
|
||||
print(f"\nCheckpoint has {len(state_dict)} tensors")
|
||||
total_params = sum(p.numel() for p in state_dict.values())
|
||||
print(f"Total parameters: {total_params:,} ({total_params / 1e9:.2f}B)")
|
||||
|
||||
# Convert weights (memory-efficient mode)
|
||||
print(f"\nConverting weights to {args.dtype}...")
|
||||
converted, unmapped, skipped = convert_weights(
|
||||
state_dict,
|
||||
verbose=args.verbose,
|
||||
dtype=target_dtype,
|
||||
memory_efficient=True,
|
||||
)
|
||||
|
||||
# Force garbage collection after conversion
|
||||
del state_dict
|
||||
gc.collect()
|
||||
|
||||
print(f"\nConversion summary:")
|
||||
print(f" Converted: {len(converted)} tensors")
|
||||
print(f" Skipped: {len(skipped)} tensors (dynamic/metadata)")
|
||||
print(f" Unmapped: {len(unmapped)} tensors (kept original names)")
|
||||
|
||||
if unmapped:
|
||||
print("\nWarning: The following keys were not mapped:")
|
||||
for key in unmapped[:20]:
|
||||
print(f" {key}")
|
||||
if len(unmapped) > 20:
|
||||
print(f" ... and {len(unmapped) - 20} more")
|
||||
|
||||
# Write output
|
||||
output_dir = Path(args.output)
|
||||
if output_dir.exists() and not args.force:
|
||||
raise FileExistsError(f"Output directory exists: {output_dir}. Use --force to overwrite.")
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Write transformer weights
|
||||
print(f"\nSaving converted weights...")
|
||||
transformer_config = build_transformer_config()
|
||||
write_component(output_dir, "transformer", converted, transformer_config)
|
||||
|
||||
# Free memory after saving
|
||||
del converted
|
||||
gc.collect()
|
||||
|
||||
# Write model_index.json
|
||||
model_index = build_model_index()
|
||||
model_index_path = output_dir / "model_index.json"
|
||||
with model_index_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(model_index, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"\nSaved model_index.json to {model_index_path}")
|
||||
|
||||
if args.components_source is not None:
|
||||
print("\nAdding inference components...")
|
||||
component_source_dir = resolve_model_dir(
|
||||
args.components_source,
|
||||
cache_dir=Path(args.components_cache_dir)
|
||||
if args.components_cache_dir else None,
|
||||
)
|
||||
add_inference_components(
|
||||
source_dir=component_source_dir,
|
||||
output_dir=output_dir,
|
||||
link_components=args.link_components,
|
||||
)
|
||||
patch_gen3c_vae_config(output_dir)
|
||||
else:
|
||||
print("\nNote: Only transformer/model_index were written.")
|
||||
print("FastVideo local loading also requires: vae/, text_encoder/, tokenizer/, scheduler/.")
|
||||
print("Re-run with --components-source to add them automatically.")
|
||||
|
||||
print(f"\nConversion complete! Output saved to {output_dir}")
|
||||
print("\nTo use with FastVideo:")
|
||||
print(f" from fastvideo.models.dits.gen3c import Gen3CTransformer3DModel")
|
||||
print(f" model = Gen3CTransformer3DModel.from_pretrained('{output_dir}/transformer')")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# Inference Configs
|
||||
|
||||
These files are nested inference configs for the config-first CLI.
|
||||
|
||||
Run them with:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config scripts/inference/<config>.yaml
|
||||
```
|
||||
|
||||
Or use the helper wrapper:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/run.sh scripts/inference/<config>.yaml
|
||||
```
|
||||
|
||||
Override config values with dotted paths:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config scripts/inference/<config>.yaml \
|
||||
--request.sampling.seed 42 \
|
||||
--request.prompt "A panda skiing at sunset"
|
||||
```
|
||||
|
||||
The same overrides work through the wrapper:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/run.sh scripts/inference/<config>.yaml \
|
||||
--generator.engine.num_gpus 2 \
|
||||
--request.output.output_path outputs/custom_run
|
||||
```
|
||||
|
||||
Some configs require an attention backend environment variable. When needed,
|
||||
the file header shows the exact command to use.
|
||||
@@ -0,0 +1,24 @@
|
||||
# Run with:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN fastvideo generate --config scripts/inference/inference_fasthunyuan.yaml
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-Diffusers
|
||||
engine:
|
||||
num_gpus: 4
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 4
|
||||
pipeline:
|
||||
experimental:
|
||||
embedded_cfg_scale: 6
|
||||
flow_shift: 17
|
||||
request:
|
||||
prompt: A beautiful woman in a red dress walking down a street
|
||||
sampling:
|
||||
seed: 1024
|
||||
num_frames: 125
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
guidance_scale: 1
|
||||
output:
|
||||
output_path: outputs_video/
|
||||
@@ -0,0 +1,24 @@
|
||||
# Run with:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN fastvideo generate --config scripts/inference/inference_hunyuan.yaml
|
||||
generator:
|
||||
model_path: hunyuanvideo-community/HunyuanVideo
|
||||
engine:
|
||||
num_gpus: 4
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 4
|
||||
pipeline:
|
||||
experimental:
|
||||
embedded_cfg_scale: 6
|
||||
flow_shift: 7
|
||||
request:
|
||||
prompt: A beautiful woman in a red dress walking down a street
|
||||
sampling:
|
||||
seed: 1024
|
||||
num_frames: 125
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 1
|
||||
output:
|
||||
output_path: outputs_video/
|
||||
@@ -0,0 +1,43 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_longcat.yaml
|
||||
generator:
|
||||
model_path: FastVideo/LongCat-Video-T2V-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
text_encoder: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
experimental:
|
||||
enable_bsa: false
|
||||
request:
|
||||
prompt: >-
|
||||
In a realistic photography style, a white boy around seven or eight years
|
||||
old sits on a park bench, wearing a light blue T-shirt, denim shorts, and
|
||||
white sneakers. He holds an ice cream cone with vanilla and chocolate
|
||||
flavors, and beside him is a medium-sized golden Labrador. Smiling, the
|
||||
boy offers the ice cream to the dog, who eagerly licks it with its tongue.
|
||||
The sun is shining brightly, and the background features a green lawn and
|
||||
several tall trees, creating a warm and loving scene.
|
||||
negative_prompt: >-
|
||||
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
|
||||
sampling:
|
||||
seed: 42
|
||||
num_frames: 93
|
||||
height: 480
|
||||
width: 832
|
||||
fps: 15
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 4.0
|
||||
output:
|
||||
output_path: outputs_video/longcat_t2v
|
||||
@@ -0,0 +1,44 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_longcat_distill.yaml
|
||||
generator:
|
||||
model_path: FastVideo/LongCat-Video-T2V-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
components:
|
||||
lora_path: FastVideo/LongCat-Video-T2V-Distilled-LoRA
|
||||
experimental:
|
||||
enable_bsa: false
|
||||
lora_nickname: distilled
|
||||
request:
|
||||
prompt: >-
|
||||
In a realistic photography style, a white boy around seven or eight years
|
||||
old sits on a park bench, wearing a light blue T-shirt, denim shorts, and
|
||||
white sneakers. He holds an ice cream cone with vanilla and chocolate
|
||||
flavors, and beside him is a medium-sized golden Labrador. Smiling, the
|
||||
boy offers the ice cream to the dog, who eagerly licks it with its tongue.
|
||||
The sun is shining brightly, and the background features a green lawn and
|
||||
several tall trees, creating a warm and loving scene.
|
||||
negative_prompt: >-
|
||||
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
|
||||
sampling:
|
||||
seed: 42
|
||||
num_frames: 93
|
||||
height: 480
|
||||
width: 832
|
||||
fps: 15
|
||||
num_inference_steps: 16
|
||||
guidance_scale: 1.0
|
||||
output:
|
||||
output_path: outputs_video/longcat_distill
|
||||
@@ -0,0 +1,42 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_longcat_i2v.yaml
|
||||
generator:
|
||||
model_path: FastVideo/LongCat-Video-I2V-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
experimental:
|
||||
enable_bsa: false
|
||||
request:
|
||||
prompt: >-
|
||||
A woman sits at a wooden table by the window in a cozy café. She reaches
|
||||
out with her right hand, picks up the white coffee cup from the saucer,
|
||||
and gently brings it to her lips to take a sip. After drinking, she places
|
||||
the cup back on the table and looks out the window, enjoying the peaceful
|
||||
atmosphere.
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
image_path: assets/girl.png
|
||||
sampling:
|
||||
seed: 42
|
||||
num_frames: 93
|
||||
height: 480
|
||||
width: 480
|
||||
fps: 15
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 4.0
|
||||
output:
|
||||
output_path: outputs_video/longcat_i2v
|
||||
@@ -0,0 +1,51 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_longcat_refine_fromvideo.yaml
|
||||
generator:
|
||||
model_path: FastVideo/LongCat-Video-T2V-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
components:
|
||||
lora_path: FastVideo/LongCat-Video-T2V-Refinement-LoRA
|
||||
experimental:
|
||||
enable_bsa: true
|
||||
bsa_sparsity: 0.875
|
||||
bsa_chunk_q: [4, 4, 8]
|
||||
bsa_chunk_k: [4, 4, 8]
|
||||
lora_nickname: refinement
|
||||
request:
|
||||
prompt: >-
|
||||
In a realistic photography style, a white boy around seven or eight years
|
||||
old sits on a park bench, wearing a light blue T-shirt, denim shorts, and
|
||||
white sneakers. He holds an ice cream cone with vanilla and chocolate
|
||||
flavors, and beside him is a medium-sized golden Labrador. Smiling, the
|
||||
boy offers the ice cream to the dog, who eagerly licks it with its tongue.
|
||||
The sun is shining brightly, and the background features a green lawn and
|
||||
several tall trees, creating a warm and loving scene.
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
refine_from: outputs_video/longcat_distill/In a realistic photography style, a white boy around seven or eight years old sits on a park bench,.mp4
|
||||
sampling:
|
||||
seed: 42
|
||||
height: 720
|
||||
width: 1280
|
||||
fps: 30
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 1.0
|
||||
extensions:
|
||||
t_thresh: 0.5
|
||||
spatial_refine_only: false
|
||||
num_cond_frames: 0
|
||||
output:
|
||||
output_path: outputs_video/longcat_refine_720p
|
||||
@@ -0,0 +1,43 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_longcat_vc.yaml
|
||||
generator:
|
||||
model_path: FastVideo/LongCat-Video-VC-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
experimental:
|
||||
enable_bsa: false
|
||||
request:
|
||||
prompt: >-
|
||||
A person rides a motorcycle along a long, straight road that stretches
|
||||
between a body of water and a forested hillside. The rider steadily
|
||||
accelerates, keeping the motorcycle centered between the guardrails, while
|
||||
the scenery passes by on both sides. The video captures the journey from
|
||||
the rider's perspective, emphasizing the sense of motion and adventure.
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
video_path: assets/motorcycle.mp4
|
||||
sampling:
|
||||
seed: 42
|
||||
num_frames: 93
|
||||
height: 480
|
||||
width: 832
|
||||
fps: 15
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 4.0
|
||||
extensions:
|
||||
num_cond_frames: 13
|
||||
output:
|
||||
output_path: outputs_video/longcat_vc
|
||||
@@ -0,0 +1,36 @@
|
||||
# Run with:
|
||||
# fastvideo generate --config scripts/inference/inference_wan.yaml
|
||||
generator:
|
||||
model_path: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
experimental:
|
||||
flow_shift: 8.0
|
||||
request:
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
prompt_path: assets/prompt.txt
|
||||
sampling:
|
||||
seed: 1024
|
||||
num_frames: 77
|
||||
height: 480
|
||||
width: 832
|
||||
fps: 16
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 6.0
|
||||
output:
|
||||
output_path: outputs_video/
|
||||
@@ -0,0 +1,37 @@
|
||||
# Run with:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN fastvideo generate --config scripts/inference/inference_wan_1.3B_VMoba.yaml
|
||||
generator:
|
||||
model_path: FastVideo/Wan2.1-T2V-1.3B-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
experimental:
|
||||
flow_shift: 8.0
|
||||
moba_config_path: fastvideo/configs/backend/vmoba/wan_1.3B_77_480_832.json
|
||||
request:
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
prompt_path: assets/prompt.txt
|
||||
sampling:
|
||||
seed: 1024
|
||||
num_frames: 77
|
||||
height: 480
|
||||
width: 832
|
||||
fps: 16
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 6.0
|
||||
output:
|
||||
output_path: outputs_video/
|
||||
@@ -0,0 +1,37 @@
|
||||
# Run with:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA.yaml
|
||||
generator:
|
||||
model_path: FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
offload:
|
||||
dit: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
pipeline:
|
||||
experimental:
|
||||
flow_shift: 5.0
|
||||
VSA_sparsity: 0.9
|
||||
request:
|
||||
negative_prompt: >-
|
||||
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
|
||||
inputs:
|
||||
prompt_path: assets/prompt.txt
|
||||
sampling:
|
||||
seed: 1024
|
||||
num_frames: 77
|
||||
height: 448
|
||||
width: 832
|
||||
fps: 16
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 5.0
|
||||
output:
|
||||
output_path: outputs_Wan-VSA-14B/
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user