Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8fa6ba6178 | ||
|
|
2e5fef787b | ||
|
|
43d87816bd | ||
|
|
2ace7dc6f4 | ||
|
|
0ca75db738 | ||
|
|
375ffd3fd5 | ||
|
|
2615ba4291 | ||
|
|
474dd71f28 | ||
|
|
98ad2d2db6 |
+58
-97
@@ -2,46 +2,20 @@ env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
|
||||
notify:
|
||||
- github_commit_status:
|
||||
context: "fastcheck-passed"
|
||||
if: build.env("TEST_SCOPE") == "fastcheck" || build.env("TEST_SCOPE") == null
|
||||
- github_commit_status:
|
||||
context: "full-suite-passed"
|
||||
if: build.env("TEST_SCOPE") == "full"
|
||||
|
||||
steps:
|
||||
# ============================================================
|
||||
- label: ":dart: Direct Test (${TEST_TYPE})"
|
||||
if: build.env("TEST_SCOPE") == "direct"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
# ============================================================
|
||||
# Fastcheck: Runs on every PR (~10-15 min parallel)
|
||||
# Core component validation: encoders, VAEs, transformers,
|
||||
# CUDA kernels, and unit tests.
|
||||
# ============================================================
|
||||
- label: "Trigger Fastcheck"
|
||||
if: build.env("TEST_SCOPE") == "fastcheck" || build.env("TEST_SCOPE") == null
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: 'git fetch origin "${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}" && git diff --name-only "origin/${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}...HEAD"'
|
||||
watch:
|
||||
- path:
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: 'git fetch origin "$BUILDKITE_PULL_REQUEST_BASE_BRANCH" && git diff --name-only origin/"$BUILDKITE_PULL_REQUEST_BASE_BRANCH"...HEAD'
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/models/encoders/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/encoders/**"
|
||||
@@ -49,7 +23,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
label: ":microscope: Encoder Tests"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
@@ -62,7 +36,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
label: ":microscope: VAE Tests"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
@@ -77,62 +51,18 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":microscope: Transformer Tests"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo-kernel/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":microscope: Kernel Tests"
|
||||
env:
|
||||
- TEST_TYPE=kernel_tests
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- ".buildkite/**"
|
||||
- ".github/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":microscope: Unit Tests"
|
||||
env:
|
||||
- TEST_TYPE=unit_test
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
# ============================================================
|
||||
# Full Suite: Runs when TEST_SCOPE=full
|
||||
# Triggered by adding the 'ready' label (via ci-trigger-full-suite.yml)
|
||||
# or on-demand via /test full slash command.
|
||||
# Includes integration tests, SSIM regression, training pipelines,
|
||||
# and performance benchmarks.
|
||||
# ============================================================
|
||||
- label: "Trigger Full Suite"
|
||||
if: build.env("TEST_SCOPE") == "full"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: 'git fetch origin "${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}" && git diff --name-only "origin/${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}...HEAD"'
|
||||
watch:
|
||||
- path:
|
||||
- path:
|
||||
- "fastvideo/**/*.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":bar_chart: SSIM Tests"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
@@ -147,7 +77,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Inference Tests"
|
||||
label: "LoRA Inference Tests"
|
||||
env:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
@@ -158,7 +88,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Training Tests"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
@@ -169,7 +99,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Distillation DMD Tests"
|
||||
label: "Distillation DMDTests"
|
||||
env:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
agents:
|
||||
@@ -181,7 +111,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Self-Forcing Tests"
|
||||
label: "Self-Forcing Tests"
|
||||
env:
|
||||
- TEST_TYPE=self_forcing
|
||||
agents:
|
||||
@@ -192,7 +122,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
label: "LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
agents:
|
||||
@@ -204,11 +134,20 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Training Tests VSA"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo-kernel/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Kernel Tests"
|
||||
env:
|
||||
- TEST_TYPE=kernel_tests
|
||||
- path:
|
||||
- "fastvideo-kernel/**"
|
||||
- "fastvideo/attention/backends/vmoba.py"
|
||||
@@ -216,11 +155,22 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Inference Tests VMoBA"
|
||||
env:
|
||||
label: "Inference Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Unit Tests"
|
||||
env:
|
||||
- TEST_TYPE=unit_test
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/models/dits/**"
|
||||
- "fastvideo/pipelines/**"
|
||||
@@ -234,7 +184,7 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Performance Tests"
|
||||
label: "Performance Tests"
|
||||
env:
|
||||
- TEST_TYPE=performance
|
||||
agents:
|
||||
@@ -247,8 +197,19 @@ steps:
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: API Server Tests"
|
||||
label: "API Server Tests"
|
||||
env:
|
||||
- TEST_TYPE=api_server
|
||||
agents:
|
||||
queue: "default"
|
||||
# - path:
|
||||
# - "scripts/lora_extraction/**"
|
||||
# - "pyproject.toml"
|
||||
# - "docker/Dockerfile.python3.12"
|
||||
# config:
|
||||
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
# label: "LoRA Extraction Tests"
|
||||
# env:
|
||||
# - TEST_TYPE=lora_extraction
|
||||
# agents:
|
||||
# queue: "default"
|
||||
|
||||
@@ -59,11 +59,7 @@ if [ -z "${TEST_TYPE:-}" ]; then
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
|
||||
@@ -1,22 +1,3 @@
|
||||
<!--
|
||||
PR TITLE: Must start with a type tag, e.g.:
|
||||
[feat] Add new model [bugfix] Fix VAE tiling [refactor] Restructure pipeline
|
||||
[perf] Optimize kernel [ci] Update tests [docs] Add guide
|
||||
[misc] Cleanup configs [new-model] Port Flux2
|
||||
|
||||
MERGE WORKFLOW:
|
||||
1. Ensure pre-commit passes and you have at least 1 approval
|
||||
2. Comment /merge (or add the "ready" label) to enter the Merge Queue
|
||||
3. Full Test Suite runs automatically on a staging branch → auto-merge on success
|
||||
|
||||
ON-DEMAND TESTING (write access required):
|
||||
/test full — Full Test Suite /test ssim — SSIM regression
|
||||
/test training — Training pipeline /test encoder — Encoder tests
|
||||
/test transformer — Transformer tests /test vae — VAE tests
|
||||
/test kernel — CUDA kernel tests /test unit — Unit tests
|
||||
See docs/contributing/pull_requests.md for all 17 test commands
|
||||
-->
|
||||
|
||||
## Purpose
|
||||
|
||||
<!-- What does this PR do? Link the related issue if applicable. -->
|
||||
|
||||
@@ -1,321 +0,0 @@
|
||||
merge_protections:
|
||||
- name: PR merge requirements
|
||||
if:
|
||||
- base = main
|
||||
success_conditions:
|
||||
- "title~=(?i)^\\[(feat|feature|bugfix|fix|refactor|perf|ci|doc|docs|misc|chore|kernel|new.?model)\\]"
|
||||
- check-success~=pre-commit
|
||||
- check-success=fastcheck-passed
|
||||
|
||||
pull_request_rules:
|
||||
|
||||
# ============================================================
|
||||
# Type labels (from PR title prefix)
|
||||
# ============================================================
|
||||
|
||||
- name: "label type: feat"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[(feat|feature)\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: feat"]
|
||||
|
||||
- name: "label type: bugfix"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[(bug)?fix\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: bugfix"]
|
||||
|
||||
- name: "label type: refactor"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[refactor\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: refactor"]
|
||||
|
||||
- name: "label type: perf"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[perf\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: perf"]
|
||||
|
||||
- name: "label type: ci"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[ci\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: ci"]
|
||||
|
||||
- name: "label type: docs"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[(doc|docs)\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: docs"]
|
||||
|
||||
- name: "label type: misc"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[(misc|chore)\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: misc"]
|
||||
|
||||
- name: "label type: new-model"
|
||||
conditions:
|
||||
- "title~=(?i)^\\[new.?model\\]"
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["type: new-model"]
|
||||
|
||||
# ============================================================
|
||||
# Scope labels (from changed files)
|
||||
# ============================================================
|
||||
|
||||
- name: "label scope: training"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^fastvideo/train/
|
||||
- files~=^fastvideo/training/
|
||||
- files~=^fastvideo/distillation/
|
||||
- files~=^examples/train/
|
||||
- files~=^examples/training/
|
||||
- files~=^examples/distill/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: training"]
|
||||
|
||||
- name: "label scope: inference"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^fastvideo/pipelines/basic/
|
||||
- files~=^fastvideo/pipelines/stages/
|
||||
- files~=^fastvideo/pipelines/samplers/
|
||||
- files~=^fastvideo/entrypoints/
|
||||
- files~=^fastvideo/worker/
|
||||
- files~=^fastvideo/configs/sample/
|
||||
- files~=^fastvideo/configs/pipelines/
|
||||
- files~=^examples/inference/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: inference"]
|
||||
|
||||
- name: "label scope: attention"
|
||||
conditions:
|
||||
- files~=^fastvideo/attention/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: attention"]
|
||||
|
||||
- name: "label scope: kernel"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^fastvideo-kernel/
|
||||
- files~=^csrc/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: kernel"]
|
||||
|
||||
- name: "label scope: data"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^fastvideo/dataset/
|
||||
- files~=^fastvideo/pipelines/preprocess/
|
||||
- files~=^examples/preprocessing/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: data"]
|
||||
|
||||
- name: "label scope: infra"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^\.github/
|
||||
- files~=^\.buildkite/
|
||||
- files~=^fastvideo/tests/
|
||||
- files~=^docker/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: infra"]
|
||||
|
||||
- name: "label scope: distributed"
|
||||
conditions:
|
||||
- files~=^fastvideo/distributed/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: distributed"]
|
||||
|
||||
- name: "label scope: docs"
|
||||
conditions:
|
||||
- files~=^docs/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: docs"]
|
||||
|
||||
- name: "label scope: ui"
|
||||
conditions:
|
||||
- files~=^ui/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: ui"]
|
||||
|
||||
- name: "label scope: model"
|
||||
conditions:
|
||||
- or:
|
||||
- files~=^fastvideo/models/
|
||||
- files~=^fastvideo/layers/
|
||||
- files~=^fastvideo/configs/models/
|
||||
- -closed
|
||||
actions:
|
||||
label:
|
||||
add: ["scope: model"]
|
||||
|
||||
# ============================================================
|
||||
# Pre-commit failure help comment
|
||||
# ============================================================
|
||||
|
||||
- name: comment on pre-commit failure
|
||||
conditions:
|
||||
- check-failure~=pre-commit
|
||||
- -closed
|
||||
actions:
|
||||
comment:
|
||||
message: |
|
||||
## Pre-commit checks failed
|
||||
|
||||
Hi @{{author}}, the pre-commit checks have failed. To fix them locally:
|
||||
|
||||
```bash
|
||||
# Install pre-commit if you haven't already
|
||||
uv pip install pre-commit
|
||||
pre-commit install
|
||||
|
||||
# Run all checks and auto-fix what's possible
|
||||
pre-commit run --all-files
|
||||
```
|
||||
|
||||
Common fixes:
|
||||
- **yapf**: `yapf -i <file>` (formatting)
|
||||
- **ruff**: `ruff check --fix <file>` (linting)
|
||||
- **codespell**: `codespell --write-changes <file>` (spelling)
|
||||
|
||||
After fixing, commit and push the changes. The checks will re-run automatically.
|
||||
|
||||
For future commits, `pre-commit` will run automatically on changed files before each commit.
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Merge conflict detection
|
||||
# ============================================================
|
||||
|
||||
- name: label conflicting PRs
|
||||
conditions:
|
||||
- conflict
|
||||
- -closed
|
||||
- label!=stale
|
||||
actions:
|
||||
label:
|
||||
add: [needs-rebase]
|
||||
comment:
|
||||
message: |
|
||||
This PR has merge conflicts with the base branch. Please rebase:
|
||||
|
||||
```bash
|
||||
git fetch origin main
|
||||
git rebase origin/main
|
||||
# Resolve any conflicts, then:
|
||||
git push --force-with-lease
|
||||
```
|
||||
|
||||
- name: remove conflict label when resolved
|
||||
conditions:
|
||||
- -conflict
|
||||
- -closed
|
||||
- label=needs-rebase
|
||||
actions:
|
||||
label:
|
||||
remove: [needs-rebase]
|
||||
|
||||
# ============================================================
|
||||
# Auto-merge and auto-rebase
|
||||
# ============================================================
|
||||
|
||||
- name: auto-merge when ready and all checks pass
|
||||
conditions:
|
||||
- label=ready
|
||||
- "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
|
||||
- -conflict
|
||||
- -closed
|
||||
- -draft
|
||||
actions:
|
||||
merge:
|
||||
method: squash
|
||||
|
||||
- name: auto-rebase when ready and Full Suite passed
|
||||
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]
|
||||
|
||||
# ============================================================
|
||||
# PR title format help
|
||||
# ============================================================
|
||||
|
||||
- name: comment on invalid PR title format
|
||||
conditions:
|
||||
- -closed
|
||||
- -draft
|
||||
- "-title~=(?i)^\\[(feat|feature|bugfix|fix|refactor|perf|ci|doc|docs|misc|chore|kernel|new.?model)\\]"
|
||||
actions:
|
||||
comment:
|
||||
message: |
|
||||
## ⚠️ PR title format required
|
||||
|
||||
Your PR title must start with a type tag in brackets. Examples:
|
||||
- `[feat] Add new model support`
|
||||
- `[bugfix] Fix VAE tiling corruption`
|
||||
- `[refactor] Restructure training pipeline`
|
||||
- `[perf] Optimize attention kernel`
|
||||
- `[ci] Update test infrastructure`
|
||||
- `[docs] Add inference guide`
|
||||
- `[misc] Clean up configs`
|
||||
- `[new-model] Port Flux2 to FastVideo`
|
||||
|
||||
Valid tags: `feat`, `feature`, `bugfix`, `fix`, `refactor`, `perf`, `ci`, `doc`, `docs`, `misc`, `chore`, `kernel`, `new-model`
|
||||
|
||||
Please update your PR title and the merge protection check will pass automatically.
|
||||
|
||||
@@ -0,0 +1,249 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def parse_arguments():
|
||||
"""Parse command line arguments"""
|
||||
parser = argparse.ArgumentParser(description='Run tests on RunPod GPU')
|
||||
parser.add_argument('--gpu-type', type=str, help='GPU type to use')
|
||||
parser.add_argument('--gpu-count',
|
||||
type=int,
|
||||
help='Number of GPUs to use',
|
||||
default=1)
|
||||
parser.add_argument('--test-command', type=str, help='Test command to run')
|
||||
parser.add_argument('--disk-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Container disk size in GB (default: 20)')
|
||||
parser.add_argument('--volume-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Persistent volume size in GB (default: 20)')
|
||||
parser.add_argument(
|
||||
'--image',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Docker image to use')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
args = parse_arguments()
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ['GITHUB_RUN_ID']
|
||||
JOB_ID = os.environ['JOB_ID']
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def create_pod():
|
||||
"""Create a RunPod instance"""
|
||||
# Ensure image name is lowercase (Docker requirement)
|
||||
image_name = args.image.lower()
|
||||
print(f"Using specified image: {image_name}")
|
||||
|
||||
docker_start_cmd = [
|
||||
"bash",
|
||||
"-c",
|
||||
"apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
]
|
||||
|
||||
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
|
||||
payload = {
|
||||
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
|
||||
"containerDiskInGb": args.disk_size,
|
||||
"volumeInGb": args.volume_size,
|
||||
"gpuTypeIds": [args.gpu_type],
|
||||
"gpuCount": args.gpu_count,
|
||||
"imageName": image_name,
|
||||
"allowedCudaVersions": ["12.4"],
|
||||
"dockerStartCmd": docker_start_cmd
|
||||
}
|
||||
|
||||
response = requests.post(PODS_API, headers=HEADERS, json=payload)
|
||||
response_data = response.json()
|
||||
print(f"Response: {json.dumps(response_data, indent=2)}")
|
||||
|
||||
return response_data["id"]
|
||||
|
||||
|
||||
def wait_for_pod(pod_id):
|
||||
"""Wait for pod to be in RUNNING state and fully ready with SSH access"""
|
||||
print("Waiting for RunPod to be ready...")
|
||||
|
||||
# First wait for RUNNING status
|
||||
max_attempts = 10
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
status = pod_data["desiredStatus"]
|
||||
|
||||
if status == "RUNNING":
|
||||
print("RunPod is running! Now waiting for ports to be assigned...")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Current status: {status}, waiting... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(2)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError(
|
||||
"Timed out waiting for RunPod to reach RUNNING state")
|
||||
|
||||
# Wait for ports to be assigned
|
||||
max_attempts = 50
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
port_mappings = pod_data.get("portMappings")
|
||||
|
||||
if (port_mappings is not None and "22" in port_mappings
|
||||
and pod_data.get("publicIp", "") != ""):
|
||||
print("RunPod is ready with SSH access!")
|
||||
print(f"SSH IP: {pod_data['publicIp']}")
|
||||
print(f"SSH Port: {port_mappings['22']}")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(20)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError("Timed out waiting for RunPod SSH access")
|
||||
|
||||
|
||||
def execute_command(pod_id):
|
||||
"""Execute command on the pod via SSH using system SSH client"""
|
||||
print(f"Running command: {args.test_command}")
|
||||
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
ssh_ip = pod_data["publicIp"]
|
||||
ssh_port = pod_data["portMappings"]["22"]
|
||||
|
||||
# Copy the repository to the pod using scp
|
||||
repo_dir = os.path.abspath(os.getcwd())
|
||||
repo_name = os.path.basename(repo_dir)
|
||||
|
||||
print(f"Copying repository from {repo_dir} to RunPod...")
|
||||
|
||||
tar_command = [
|
||||
"tar", "-czf", "/tmp/repo.tar.gz", "-C",
|
||||
os.path.dirname(repo_dir), repo_name
|
||||
]
|
||||
subprocess.run(tar_command, check=True)
|
||||
|
||||
# Copy the tarball to the pod
|
||||
scp_command = [
|
||||
"scp", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-P",
|
||||
str(ssh_port), "/tmp/repo.tar.gz", f"root@{ssh_ip}:/tmp/"
|
||||
]
|
||||
subprocess.run(scp_command, check=True)
|
||||
|
||||
# For custom image, we can use the pre-configured environment
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
remote_command = " && ".join(setup_steps)
|
||||
|
||||
ssh_command = [
|
||||
"ssh", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-p",
|
||||
str(ssh_port), f"root@{ssh_ip}", remote_command
|
||||
]
|
||||
|
||||
print(f"Connecting to {ssh_ip}:{ssh_port}...")
|
||||
|
||||
try:
|
||||
process = subprocess.Popen(ssh_command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=0)
|
||||
|
||||
stdout_lines = []
|
||||
|
||||
print("Command output:")
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line.strip())
|
||||
stdout_lines.append(line)
|
||||
|
||||
process.wait()
|
||||
|
||||
return_code = process.returncode
|
||||
success = return_code == 0
|
||||
|
||||
stdout_str = "".join(stdout_lines)
|
||||
|
||||
if success:
|
||||
print("Command executed successfully")
|
||||
else:
|
||||
print(f"Command failed with exit code {return_code}")
|
||||
|
||||
result = {
|
||||
"success": success,
|
||||
"return_code": return_code,
|
||||
"stdout": stdout_str,
|
||||
"stderr": ""
|
||||
}
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error executing SSH command: {str(e)}")
|
||||
result = {"success": False, "error": str(e), "stdout": "", "stderr": ""}
|
||||
return result
|
||||
|
||||
|
||||
def terminate_pod(pod_id):
|
||||
"""Terminate the pod"""
|
||||
print("Terminating RunPod...")
|
||||
requests.delete(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
print(f"Terminated pod {pod_id}")
|
||||
|
||||
|
||||
def main():
|
||||
pod_id = None
|
||||
try:
|
||||
pod_id = create_pod()
|
||||
wait_for_pod(pod_id)
|
||||
result = execute_command(pod_id)
|
||||
|
||||
if result.get("error") is not None:
|
||||
print(f"Error executing command: {result['error']}")
|
||||
sys.exit(1)
|
||||
|
||||
if not result.get("success", False):
|
||||
print(
|
||||
"Tests failed - check the output above for details on which tests failed"
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
finally:
|
||||
if pod_id:
|
||||
terminate_pod(pod_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,90 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ.get('GITHUB_RUN_ID', str(uuid.uuid4()))
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def get_job_ids():
|
||||
"""Parse job IDs from environment variable"""
|
||||
job_ids_str = os.environ.get('JOB_IDS')
|
||||
try:
|
||||
job_ids = json.loads(job_ids_str)
|
||||
if not isinstance(job_ids, list):
|
||||
print("Error: JOB_IDS is not a list.")
|
||||
sys.exit(1)
|
||||
return job_ids
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JOB_IDS: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cleanup_pods():
|
||||
"""Find and terminate RunPod instances"""
|
||||
print(f"Run ID: {RUN_ID}")
|
||||
|
||||
single_job_id = os.environ.get('JOB_ID')
|
||||
|
||||
if single_job_id:
|
||||
job_ids = [single_job_id]
|
||||
print(f"Job ID: {single_job_id}")
|
||||
else:
|
||||
job_ids = get_job_ids()
|
||||
print(f"Job IDs: {job_ids}")
|
||||
|
||||
# Get all pods associated with RunPod API_KEY
|
||||
try:
|
||||
response = requests.get(PODS_API, headers=HEADERS)
|
||||
response.raise_for_status()
|
||||
pods = response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error getting pods: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# Find and terminate pods created by this workflow run
|
||||
terminated_pods = []
|
||||
for pod in pods:
|
||||
pod_name = pod.get("name", "")
|
||||
pod_id = pod.get("id")
|
||||
|
||||
# Check if this pod was created by one of our jobs
|
||||
if any(f"{job_id}-{RUN_ID}" in pod_name for job_id in job_ids):
|
||||
print(f"Found pod: {pod_id} ({pod_name})")
|
||||
try:
|
||||
print(f"Terminating pod {pod_id}...")
|
||||
term_response = requests.delete(f"{PODS_API}/{pod_id}",
|
||||
headers=HEADERS)
|
||||
term_response.raise_for_status()
|
||||
terminated_pods.append(pod_id)
|
||||
print(f"Successfully terminated pod {pod_id}")
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error terminating pod {pod_id}: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
if terminated_pods:
|
||||
if single_job_id:
|
||||
print(f"Terminated pod: {terminated_pods[0]}")
|
||||
else:
|
||||
print(f"Terminated {len(terminated_pods)} pods: {terminated_pods}")
|
||||
else:
|
||||
if single_job_id:
|
||||
print(f"No pod found matching pattern: {single_job_id}-{RUN_ID}")
|
||||
else:
|
||||
print("No pods found to terminate.")
|
||||
|
||||
|
||||
def main():
|
||||
cleanup_pods()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -32,7 +32,7 @@ permissions:
|
||||
jobs:
|
||||
build-python-3-10:
|
||||
if: ${{ github.event.inputs.python_3_10 == 'true' }}
|
||||
uses: ./.github/workflows/_template-build-image.yml
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.10'
|
||||
dockerfile_path: docker/Dockerfile.python3.10
|
||||
@@ -41,7 +41,7 @@ jobs:
|
||||
|
||||
build-python-3-11:
|
||||
if: ${{ github.event.inputs.python_3_11 == 'true' }}
|
||||
uses: ./.github/workflows/_template-build-image.yml
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.11'
|
||||
dockerfile_path: docker/Dockerfile.python3.11
|
||||
@@ -50,7 +50,7 @@ jobs:
|
||||
|
||||
build-python-3-12:
|
||||
if: ${{ github.event.inputs.python_3_12 == 'true' }}
|
||||
uses: ./.github/workflows/_template-build-image.yml
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12
|
||||
@@ -59,9 +59,9 @@ jobs:
|
||||
|
||||
build-python-3-12-cuda-12-9:
|
||||
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
|
||||
uses: ./.github/workflows/_template-build-image.yml
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
|
||||
tag_suffix: py3.12-cuda12.9.1
|
||||
secrets: inherit
|
||||
secrets: inherit
|
||||
@@ -1,29 +0,0 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
|
||||
concurrency:
|
||||
group: pre-commit-${{ github.ref }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/ruff.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
@@ -1,230 +0,0 @@
|
||||
name: Slash Commands
|
||||
|
||||
on:
|
||||
issue_comment:
|
||||
types: [created]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
statuses: write
|
||||
|
||||
jobs:
|
||||
handle-merge:
|
||||
if: >-
|
||||
github.event.issue.pull_request != null
|
||||
&& startsWith(github.event.comment.body, '/merge')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check write permission
|
||||
id: perm
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const { data: perm } = await github.rest.repos.getCollaboratorPermissionLevel({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
username: context.payload.comment.user.login,
|
||||
});
|
||||
const hasWrite = ['admin', 'write'].includes(perm.permission);
|
||||
if (!hasWrite) {
|
||||
core.setFailed(`User ${context.payload.comment.user.login} lacks write permission (has: ${perm.permission}).`);
|
||||
}
|
||||
core.setOutput('has_write', String(hasWrite));
|
||||
|
||||
- name: Add ready label and react
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
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({
|
||||
owner, repo,
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
parse-command:
|
||||
if: >-
|
||||
github.event.issue.pull_request != null
|
||||
&& startsWith(github.event.comment.body, '/test')
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
test_type: ${{ steps.parse.outputs.test_type }}
|
||||
test_scope: ${{ steps.parse.outputs.test_scope }}
|
||||
full_suite: ${{ steps.parse.outputs.full_suite }}
|
||||
pr_sha: ${{ steps.pr.outputs.sha }}
|
||||
pr_branch: ${{ steps.pr.outputs.branch }}
|
||||
has_write: ${{ steps.perm.outputs.has_write }}
|
||||
steps:
|
||||
- name: Check write permission
|
||||
id: perm
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const { data: perm } = await github.rest.repos.getCollaboratorPermissionLevel({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
username: context.payload.comment.user.login,
|
||||
});
|
||||
const hasWrite = ['admin', 'write'].includes(perm.permission);
|
||||
core.setOutput('has_write', String(hasWrite));
|
||||
if (!hasWrite) {
|
||||
core.info(`User ${context.payload.comment.user.login} lacks write permission — ignoring.`);
|
||||
}
|
||||
|
||||
- name: Parse /test command
|
||||
id: parse
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
shell: bash
|
||||
env:
|
||||
COMMENT: ${{ github.event.comment.body }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
declare -A MAP=(
|
||||
[encoder]=encoder [vae]=vae [transformer]=transformer
|
||||
[kernel]=kernel_tests [unit]=unit_test
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
)
|
||||
|
||||
if [ "$TEST_NAME" = "full" ]; then
|
||||
{
|
||||
echo "test_type=all"
|
||||
echo "test_scope=full"
|
||||
echo "full_suite=true"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
elif [ "$TEST_NAME" = "fastcheck" ]; then
|
||||
{
|
||||
echo "test_type=fastcheck"
|
||||
echo "test_scope=fastcheck"
|
||||
echo "full_suite=false"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
elif [ "$TEST_NAME" = "pre-commit" ]; then
|
||||
{
|
||||
echo "test_type="
|
||||
echo "test_scope=precommit"
|
||||
echo "full_suite=false"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
{
|
||||
echo "test_type=${MAP[$TEST_NAME]}"
|
||||
echo "test_scope=direct"
|
||||
echo "full_suite=false"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Get PR details
|
||||
id: pr
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const { data: pr } = await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.payload.issue.number,
|
||||
});
|
||||
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
|
||||
|
||||
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: React to comment
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
|
||||
- name: Trigger Buildkite
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_SHA: ${{ needs.parse-command.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ needs.parse-command.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ github.event.issue.number }}
|
||||
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
|
||||
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
|
||||
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
|
||||
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 "/test ${TEST_TYPE} on PR #${PR_NUMBER}" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
--arg test_scope "$TEST_SCOPE" \
|
||||
--arg full_suite "$FULL_SUITE" \
|
||||
--arg test_type "$TEST_TYPE" \
|
||||
--arg pr_number "$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: $test_scope,
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number
|
||||
}
|
||||
}')"
|
||||
@@ -1,83 +0,0 @@
|
||||
name: Trigger Full Suite
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
types: [labeled, synchronize]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
trigger:
|
||||
if: >-
|
||||
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
|| github.event.action == 'synchronize'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check ready label
|
||||
id: check
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const { data: pr } = await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.payload.pull_request.number,
|
||||
});
|
||||
const hasReady = pr.labels.some(l => l.name === 'ready');
|
||||
core.setOutput('has_ready', String(hasReady));
|
||||
if (!hasReady) core.info('No ready label — skipping Full Suite trigger.');
|
||||
|
||||
- name: Cancel previous Buildkite builds
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
run: |
|
||||
# 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')
|
||||
for build_num in $builds; do
|
||||
echo "Cancelling Buildkite build #$build_num"
|
||||
curl -sS -X PUT -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
|
||||
done
|
||||
|
||||
- name: Trigger Buildkite Full Suite
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.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}" \
|
||||
--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)
|
||||
}
|
||||
}')"
|
||||
@@ -1,65 +0,0 @@
|
||||
name: Auto-Label Issues
|
||||
|
||||
on:
|
||||
issues:
|
||||
types: [opened, edited]
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
label-issues:
|
||||
if: github.repository == 'hao-ai-lab/FastVideo'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Label by keywords
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const title = context.payload.issue.title.toLowerCase();
|
||||
const body = (context.payload.issue.body || '').toLowerCase();
|
||||
const text = title + ' ' + body;
|
||||
const labels = [];
|
||||
|
||||
const rules = [
|
||||
// scope labels (shared with PR labeling via Mergify)
|
||||
// Mapping: label → repo directories
|
||||
// scope: training → fastvideo/train/, fastvideo/training/, fastvideo/distillation/
|
||||
// scope: inference → fastvideo/pipelines/, fastvideo/entrypoints/, fastvideo/worker/
|
||||
// scope: attention → fastvideo/attention/
|
||||
// scope: kernel → fastvideo-kernel/, csrc/
|
||||
// scope: model → fastvideo/models/, fastvideo/layers/, fastvideo/configs/models/
|
||||
// scope: data → fastvideo/dataset/, fastvideo/pipelines/preprocess/
|
||||
// scope: distributed → fastvideo/distributed/
|
||||
// scope: docs → docs/
|
||||
{ keywords: ['training', 'finetune', 'fine-tune', 'lora', 'fsdp', 'distill'], label: 'scope: training' },
|
||||
{ keywords: ['inference', 'generate', 'pipeline', 'slow', 'latency'], label: 'scope: inference' },
|
||||
{ keywords: ['attention', 'vsa', 'flash', 'sta', 'vmoba', 'sparse attn'], label: 'scope: attention' },
|
||||
{ keywords: ['kernel', 'csrc', 'cuda kernel', 'thunderkittens'], label: 'scope: kernel' },
|
||||
{ keywords: ['wan', 'hunyuan', 'mochi', 'ltx', 'cogvideo', 'flux', 'sd3', 'cosmos'], label: 'scope: model' },
|
||||
{ keywords: ['dataset', 'dataloader', 'preprocessing', 'preprocess'], label: 'scope: data' },
|
||||
{ keywords: ['distributed', 'sequence parallel', 'fsdp', 'tensor parallel', 'multi-node', 'multi-gpu'], label: 'scope: distributed' },
|
||||
{ keywords: ['docs', 'documentation', 'tutorial', 'example'], label: 'scope: docs' },
|
||||
// issue-only labels (cross-module, no single repo directory)
|
||||
{ keywords: ['install', 'setup', 'pip', 'cuda', 'uv ', 'import error', 'modulenotfound'], label: 'installation' },
|
||||
{ keywords: ['memory', 'oom', 'out of memory', 'gpu memory', 'vram'], label: 'performance' },
|
||||
{ keywords: ['windows', 'macos', 'mac os', 'apple', 'mps', 'rocm', 'amd', 'npu'], label: 'platform' },
|
||||
];
|
||||
|
||||
for (const rule of rules) {
|
||||
if (rule.keywords.some(kw => text.includes(kw))) {
|
||||
labels.push(rule.label);
|
||||
}
|
||||
}
|
||||
|
||||
if (labels.length > 0) {
|
||||
await github.rest.issues.addLabels({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: context.payload.issue.number,
|
||||
labels: labels,
|
||||
});
|
||||
console.log(`Added labels: ${labels.join(', ')}`);
|
||||
} else {
|
||||
console.log('No keyword matches found');
|
||||
}
|
||||
@@ -1,51 +0,0 @@
|
||||
name: Close Stale Issues and PRs
|
||||
|
||||
on:
|
||||
schedule:
|
||||
# Daily at 1:30 AM UTC
|
||||
- cron: '30 1 * * *'
|
||||
|
||||
jobs:
|
||||
stale:
|
||||
if: github.repository == 'hao-ai-lab/FastVideo'
|
||||
permissions:
|
||||
issues: write
|
||||
pull-requests: write
|
||||
actions: write
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/stale@997185467fa4f803885201cee163a9f38240193d # v10.1.1
|
||||
with:
|
||||
operations-per-run: 500
|
||||
|
||||
exempt-draft-pr: true
|
||||
exempt-issue-labels: 'keep-open,pinned,security,Bug,RFC'
|
||||
exempt-pr-labels: 'keep-open,pinned'
|
||||
|
||||
labels-to-add-when-unstale: 'unstale'
|
||||
labels-to-remove-when-stale: 'unstale'
|
||||
|
||||
days-before-issue-stale: 90
|
||||
days-before-issue-close: 30
|
||||
stale-issue-label: 'stale'
|
||||
stale-issue-message: >
|
||||
This issue has been automatically marked as stale because it has not
|
||||
had any activity within 90 days. It will be automatically closed if
|
||||
no further activity occurs within 30 days. Leave a comment if you
|
||||
feel this issue should remain open. Thank you!
|
||||
close-issue-message: >
|
||||
This issue has been automatically closed due to inactivity. Please
|
||||
feel free to reopen if you feel it is still relevant. Thank you!
|
||||
|
||||
days-before-pr-stale: 60
|
||||
days-before-pr-close: 14
|
||||
stale-pr-label: 'stale'
|
||||
stale-pr-message: >
|
||||
This pull request has been automatically marked as stale because it
|
||||
has not had any activity within 60 days. It will be automatically
|
||||
closed if no further activity occurs within 14 days. Leave a comment
|
||||
if you feel this pull request should remain open. Thank you!
|
||||
close-pr-message: >
|
||||
This pull request has been automatically closed due to inactivity.
|
||||
Please feel free to reopen if you intend to continue working on it.
|
||||
Thank you!
|
||||
@@ -1,56 +0,0 @@
|
||||
name: Welcome First-Time Contributors
|
||||
|
||||
on:
|
||||
issues:
|
||||
types: [opened]
|
||||
pull_request_target:
|
||||
types: [opened]
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
welcome:
|
||||
if: github.repository == 'hao-ai-lab/FastVideo'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/first-interaction@34f15e814fe48ac9312ccf29db4e74fa767cbab7 # v1.3.0
|
||||
with:
|
||||
repo-token: ${{ secrets.GITHUB_TOKEN }}
|
||||
issue-message: |
|
||||
Welcome to FastVideo! Thanks for opening your first issue.
|
||||
|
||||
To help us investigate, please include:
|
||||
- **FastVideo version**: `pip show fastvideo`
|
||||
- **GPU**: `nvidia-smi` output (GPU model, driver, CUDA version)
|
||||
- **Python version**: `python --version`
|
||||
- **OS**: e.g., Ubuntu 22.04
|
||||
|
||||
If this is a bug, a minimal reproduction script helps us fix it faster.
|
||||
|
||||
Useful links:
|
||||
- [Documentation](https://hao-ai-lab.github.io/FastVideo)
|
||||
- [Contributing Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
- [Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)
|
||||
pr-message: |
|
||||
Welcome to FastVideo! Thanks for your first pull request.
|
||||
|
||||
**How our CI works:**
|
||||
|
||||
PRs run a two-tier CI system:
|
||||
1. **Pre-commit** — formatting (yapf), linting (ruff), type checking (mypy). Runs immediately on every PR.
|
||||
2. **Fastcheck** — core GPU tests (encoders, VAEs, transformers, kernels, unit tests). Runs automatically via Buildkite on relevant file changes (~10-15 min).
|
||||
3. **Full Suite** — integration tests, training pipelines, SSIM regression. Runs only when a reviewer adds the `ready` label.
|
||||
|
||||
**Before your PR is reviewed:**
|
||||
- [ ] `pre-commit run --all-files` passes locally
|
||||
- [ ] You've added or updated tests for your changes
|
||||
- [ ] The PR description explains what and why
|
||||
|
||||
If pre-commit fails, a bot comment will explain how to fix it. Fastcheck and Full Suite results appear in the Checks section below.
|
||||
|
||||
**Useful links:**
|
||||
- [Contributing Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
- [Development Roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899)
|
||||
- [Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)
|
||||
@@ -7,14 +7,14 @@ on:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/infra-docs.yml'
|
||||
- '.github/workflows/docs.yml'
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/infra-docs.yml'
|
||||
- '.github/workflows/docs.yml'
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
+3
-3
@@ -36,11 +36,11 @@ jobs:
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> "$GITHUB_OUTPUT"
|
||||
echo "new-version=$NEW_VERSION" >> "$GITHUB_OUTPUT"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> "$GITHUB_OUTPUT"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
@@ -1,17 +0,0 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "ruff",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(.+):(\\d+):(\\d+): (\\w+) (.+)$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"column": 3,
|
||||
"code": 4,
|
||||
"message": 5
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
name: PR Test
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
pull_request:
|
||||
branches: [main]
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_vae_test:
|
||||
description: "Run vae-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_transformer_test:
|
||||
description: "Run transformer-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test_VSA:
|
||||
description: "Run training-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_VSA:
|
||||
description: "Run precision-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_unit_test:
|
||||
description: "Run unit-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
change-filter:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
unit-test: ${{ steps.filter.outputs.unit-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.10'
|
||||
- 'docker/Dockerfile.python3.11'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/video_sparse_attn/**'
|
||||
- 'csrc/attn/video_sparse_attn/tk/**'
|
||||
- 'csrc/attn/video_sparse_attn/setup.py'
|
||||
- 'csrc/attn/video_sparse_attn/config_vsa.py'
|
||||
- 'csrc/attn/video_sparse_attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
unit-test:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.encoder-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_encoder_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "encoder-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
vae-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "vae-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
transformer-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "transformer-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: [
|
||||
# {version: "3.10", tag: "latest"},
|
||||
# {version: "3.11", tag: "py3.11-latest"},
|
||||
{version: "3.12", tag: "py3.12-latest"}
|
||||
]
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "ssim-test-py${{ matrix.python-version.version }}"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 2
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
precision-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_vsa.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
unit-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "unit-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs && pytest ./fastvideo/entrypoints/ -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
# nightly-test:
|
||||
# if: >-
|
||||
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
# uses: ./.github/workflows/runpod-test.yml
|
||||
# with:
|
||||
# job_id: "nightly-test"
|
||||
# gpu_type: "NVIDIA A40"
|
||||
# gpu_count: 4
|
||||
# volume_size: 100
|
||||
# disk_size: 100
|
||||
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
# timeout_minutes: 30
|
||||
# secrets:
|
||||
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
# Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, precision-test-VSA]
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -0,0 +1,18 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
@@ -0,0 +1,94 @@
|
||||
name: RunPod Test
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
job_id:
|
||||
required: true
|
||||
type: string
|
||||
description: "Unique identifier for this test job"
|
||||
gpu_type:
|
||||
required: true
|
||||
type: string
|
||||
description: "GPU type to use (e.g. NVIDIA A40, NVIDIA L40S)"
|
||||
gpu_count:
|
||||
required: true
|
||||
type: number
|
||||
description: "Number of GPUs to use"
|
||||
volume_size:
|
||||
required: false
|
||||
type: number
|
||||
default: 20
|
||||
description: "Volume size in GB"
|
||||
disk_size:
|
||||
required: false
|
||||
type: number
|
||||
default: 20
|
||||
description: "Disk size in GB"
|
||||
image:
|
||||
required: true
|
||||
type: string
|
||||
description: "Docker image to use"
|
||||
test_command:
|
||||
required: true
|
||||
type: string
|
||||
description: "Command to run tests"
|
||||
timeout_minutes:
|
||||
required: false
|
||||
type: number
|
||||
default: 30
|
||||
description: "Timeout in minutes"
|
||||
secrets:
|
||||
RUNPOD_API_KEY:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "${{ inputs.gpu_type }}"
|
||||
--gpu-count ${{ inputs.gpu_count }}
|
||||
--volume-size ${{ inputs.volume_size }}
|
||||
--disk-size ${{ inputs.disk_size }}
|
||||
--image "${{ inputs.image }}"
|
||||
--test-command "${{ inputs.test_command }}"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -0,0 +1,31 @@
|
||||
name: Run Tests
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install torch
|
||||
pip install packaging ninja
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/attn/test
|
||||
@@ -0,0 +1,257 @@
|
||||
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
|
||||
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
# For version reference https://pytorch.org/get-started/previous-versions/
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/video_sparse_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/video_sparse_attn/dist/
|
||||
@@ -33,7 +33,6 @@ env
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
@@ -85,9 +84,6 @@ docs/distillation/examples/
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
|
||||
.claude/
|
||||
.codex/
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
|
||||
@@ -18,8 +18,10 @@ exclude: |
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
examples/.*|
|
||||
.github/workflows/publish-fastvideo.yml|
|
||||
.github/workflows/_template-build-image.yml|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml|
|
||||
.github/workflows/vsa-publish.yml|
|
||||
.github/workflows/build-image-template.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
)
|
||||
repos:
|
||||
|
||||
@@ -6,17 +6,16 @@
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/03/17`: Release Live demo: [Into the Dreamverse: Vibe Directing in FastVideo](https://dreamverse.fastvideo.org/), check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release Live demo: [Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU](https://1080p.fastvideo.org/), check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py).
|
||||
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- `2025/08/04`: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
### More News
|
||||
|
||||
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389).
|
||||
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
|
||||
+8
-3
@@ -1,10 +1,15 @@
|
||||
try:
|
||||
from .comfyui.video_generator.nodes import (NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS)
|
||||
from .comfyui.video_generator.nodes import (NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS)
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
except ImportError:
|
||||
# ComfyUI environment not available, skip comfyui imports
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
|
||||
@@ -13,7 +13,8 @@ from .fvd import (
|
||||
compute_statistics,
|
||||
FVDConfig,
|
||||
)
|
||||
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor, load_extractor)
|
||||
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor,
|
||||
load_extractor)
|
||||
from .video_utils import (
|
||||
load_video_auto,
|
||||
sample_clips_from_video,
|
||||
|
||||
+41
-11
@@ -5,11 +5,18 @@ from .fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description='Compute Fréchet Video Distance (FVD)')
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Compute Fréchet Video Distance (FVD)')
|
||||
|
||||
# Required arguments
|
||||
parser.add_argument('--real-path', type=str, required=True, help='Path to real videos')
|
||||
parser.add_argument('--gen-path', type=str, required=True, help='Path to generated videos')
|
||||
parser.add_argument('--real-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to real videos')
|
||||
parser.add_argument('--gen-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to generated videos')
|
||||
|
||||
# Extractor selection
|
||||
parser.add_argument('--extractor',
|
||||
@@ -19,19 +26,42 @@ def main() -> int:
|
||||
help='Feature extractor model to use (default: i3d)')
|
||||
|
||||
# Standard args
|
||||
parser.add_argument('--seed', type=int, default=None, help='Random seed for reproducibility')
|
||||
parser.add_argument('--seed',
|
||||
type=int,
|
||||
default=None,
|
||||
help='Random seed for reproducibility')
|
||||
parser.add_argument('--protocol',
|
||||
type=str,
|
||||
default=None,
|
||||
choices=['fvd2048_16f', 'fvd2048_128f', 'quick_test'],
|
||||
help='Use standard protocol (overrides other settings)')
|
||||
parser.add_argument('--num-videos', type=int, default=2048, help='Number of videos to use')
|
||||
parser.add_argument('--num-frames', type=int, default=16, help='Number of frames per clip')
|
||||
parser.add_argument('--clip-strategy', type=str, default='beginning', help='Clip sampling strategy')
|
||||
parser.add_argument('--batch-size', type=int, default=32, help='Batch size for feature extraction')
|
||||
parser.add_argument('--device', type=str, default='cuda', help='Device to use (cuda or cpu)')
|
||||
parser.add_argument('--cache-real-features', type=str, default=None, help='Path to cache real video features')
|
||||
parser.add_argument('--quiet', action='store_true', help='Suppress progress output')
|
||||
parser.add_argument('--num-videos',
|
||||
type=int,
|
||||
default=2048,
|
||||
help='Number of videos to use')
|
||||
parser.add_argument('--num-frames',
|
||||
type=int,
|
||||
default=16,
|
||||
help='Number of frames per clip')
|
||||
parser.add_argument('--clip-strategy',
|
||||
type=str,
|
||||
default='beginning',
|
||||
help='Clip sampling strategy')
|
||||
parser.add_argument('--batch-size',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Batch size for feature extraction')
|
||||
parser.add_argument('--device',
|
||||
type=str,
|
||||
default='cuda',
|
||||
help='Device to use (cuda or cpu)')
|
||||
parser.add_argument('--cache-real-features',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Path to cache real video features')
|
||||
parser.add_argument('--quiet',
|
||||
action='store_true',
|
||||
help='Suppress progress output')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
@@ -22,7 +22,8 @@ class BaseFeatureExtractor(ABC, nn.Module):
|
||||
|
||||
def __init__(self, device: str = 'cuda'):
|
||||
super().__init__()
|
||||
self.device = torch.device(device if torch.cuda.is_available() else 'cpu')
|
||||
self.device = torch.device(
|
||||
device if torch.cuda.is_available() else 'cpu')
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
@@ -52,7 +53,10 @@ class BaseFeatureExtractor(ABC, nn.Module):
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self, videos: torch.Tensor, batch_size: int = 32, verbose: bool = True) -> torch.Tensor:
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract features for a large tensor of videos by batching.
|
||||
"""
|
||||
@@ -61,7 +65,9 @@ class BaseFeatureExtractor(ABC, nn.Module):
|
||||
|
||||
iterator = range(0, N, batch_size)
|
||||
if verbose:
|
||||
iterator = tqdm(iterator, desc=f"Extracting features ({self.__class__.__name__})")
|
||||
iterator = tqdm(
|
||||
iterator,
|
||||
desc=f"Extracting features ({self.__class__.__name__})")
|
||||
|
||||
for i in iterator:
|
||||
batch = videos[i:i + batch_size].to(self.device)
|
||||
@@ -89,7 +95,9 @@ class I3DFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def _load_model(self) -> torch.nn.Module:
|
||||
try:
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID, filename=self.MODEL_FILENAME, cache_dir=self.cache_dir)
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
return torch.jit.load(model_path, map_location=self.device)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model: {e}") from e
|
||||
@@ -111,7 +119,10 @@ class I3DFeatureExtractor(BaseFeatureExtractor):
|
||||
# Resize to 224x224
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.reshape(B * T, C, H, W)
|
||||
videos = F.interpolate(videos, size=(224, 224), mode='bilinear', align_corners=False)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.reshape(B, T, C, 224, 224)
|
||||
|
||||
# [B, T, C, H, W] -> [B, C, T, H, W]
|
||||
@@ -120,15 +131,21 @@ class I3DFeatureExtractor(BaseFeatureExtractor):
|
||||
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
batch = self.preprocess(videos)
|
||||
# TorchScript I3D returns raw logits when return_features=True
|
||||
return self.model(batch, rescale=False, resize=False, return_features=True)
|
||||
return self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
|
||||
# 2. CLIP Extractor (Semantic/Content Quality)
|
||||
class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "openai/clip-vit-base-patch32"):
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
model_name: str = "openai/clip-vit-base-patch32"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError(
|
||||
"Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.processor = CLIPProcessor.from_pretrained(model_name)
|
||||
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
|
||||
@@ -155,7 +172,9 @@ class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
images = videos.view(B * T, C, H, W)
|
||||
|
||||
# HF Processor
|
||||
inputs = self.processor(images=images, return_tensors="pt", padding=True)
|
||||
inputs = self.processor(images=images,
|
||||
return_tensors="pt",
|
||||
padding=True)
|
||||
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
||||
|
||||
# Extract features [B*T, Dim]
|
||||
@@ -169,15 +188,24 @@ class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
# 3. VideoMAE Extractor (Structure/Motion Quality)
|
||||
class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "MCG-NJU/videomae-base"):
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
model_name: str = "MCG-NJU/videomae-base"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError(
|
||||
"Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406], device=self.device).view(1, 1, 3, 1, 1))
|
||||
self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225], device=self.device).view(1, 1, 3, 1, 1))
|
||||
self.register_buffer(
|
||||
'mean',
|
||||
torch.tensor([0.485, 0.456, 0.406],
|
||||
device=self.device).view(1, 1, 3, 1, 1))
|
||||
self.register_buffer(
|
||||
'std',
|
||||
torch.tensor([0.229, 0.224, 0.225],
|
||||
device=self.device).view(1, 1, 3, 1, 1))
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
@@ -193,7 +221,10 @@ class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
# 1. Resize to 224x224
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.view(B * T, C, H, W)
|
||||
videos = F.interpolate(videos, size=(224, 224), mode='bilinear', align_corners=False)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.view(B, T, C, 224, 224)
|
||||
|
||||
# 2. Normalize to [0, 1]
|
||||
@@ -229,4 +260,5 @@ def load_extractor(name: str, device: str = 'cuda') -> BaseFeatureExtractor:
|
||||
elif name == 'videomae':
|
||||
return VideoMAEFeatureExtractor(device)
|
||||
else:
|
||||
raise ValueError(f"Unknown extractor: {name}. Options: i3d, clip, videomae")
|
||||
raise ValueError(
|
||||
f"Unknown extractor: {name}. Options: i3d, clip, videomae")
|
||||
|
||||
+47
-26
@@ -36,7 +36,8 @@ def compute_frechet_distance(mu1: np.ndarray,
|
||||
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
print(f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
|
||||
print(
|
||||
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
|
||||
covmean = covmean.real
|
||||
|
||||
trace_product = np.trace(covmean)
|
||||
@@ -66,7 +67,8 @@ class FVDConfig:
|
||||
temporal_stride: int = 1 # For sliding window clips
|
||||
|
||||
# Data processing
|
||||
video_extensions: list[str] = field(default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
|
||||
video_extensions: list[str] = field(
|
||||
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
|
||||
support_frame_dirs: bool = True
|
||||
|
||||
# Computation
|
||||
@@ -86,17 +88,25 @@ class FVDConfig:
|
||||
@classmethod
|
||||
def fvd2048_16f(cls) -> 'FVDConfig':
|
||||
"""Standard FVD protocol: 2048 videos, 16 frames, beginning clip."""
|
||||
return cls(num_videos=2048, num_frames_per_clip=16, clip_strategy='beginning', use_streaming=True)
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def fvd2048_128f(cls) -> 'FVDConfig':
|
||||
"""Long video protocol: 2048 videos, 128 frames."""
|
||||
return cls(num_videos=2048, num_frames_per_clip=128, clip_strategy='beginning', use_streaming=True)
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=128,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def quick_test(cls) -> 'FVDConfig':
|
||||
"""Quick test config: 100 videos, 16 frames."""
|
||||
return cls(num_videos=100, num_frames_per_clip=16, clip_strategy='beginning')
|
||||
return cls(num_videos=100,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning')
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Export config to dict for logging"""
|
||||
@@ -186,10 +196,14 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
max_features = config.num_videos * config.num_clips_per_video
|
||||
|
||||
if len(features) < max_features:
|
||||
print(f"WARNING: Cache has {len(features)} features but need {max_features}")
|
||||
print(
|
||||
f"WARNING: Cache has {len(features)} features but need {max_features}"
|
||||
)
|
||||
print("Cached features insufficient - will recompute...")
|
||||
elif len(features) > max_features:
|
||||
print(f"Using {max_features} features from cache (truncated from {len(features)})")
|
||||
print(
|
||||
f"Using {max_features} features from cache (truncated from {len(features)})"
|
||||
)
|
||||
features = features[:max_features]
|
||||
return features
|
||||
else:
|
||||
@@ -201,16 +215,17 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
if isinstance(videos, (str | Path)):
|
||||
target_size = (224, 224) if config.resize_before_extraction else None
|
||||
|
||||
video_generator = load_video_clips_streaming(videos,
|
||||
num_frames=config.num_frames_per_clip,
|
||||
max_videos=config.num_videos,
|
||||
clip_strategy=config.clip_strategy,
|
||||
frame_stride=config.frame_stride,
|
||||
num_clips_per_video=config.num_clips_per_video,
|
||||
video_extensions=config.video_extensions,
|
||||
support_frame_dirs=config.support_frame_dirs,
|
||||
target_size=target_size,
|
||||
verbose=True)
|
||||
video_generator = load_video_clips_streaming(
|
||||
videos,
|
||||
num_frames=config.num_frames_per_clip,
|
||||
max_videos=config.num_videos,
|
||||
clip_strategy=config.clip_strategy,
|
||||
frame_stride=config.frame_stride,
|
||||
num_clips_per_video=config.num_clips_per_video,
|
||||
video_extensions=config.video_extensions,
|
||||
support_frame_dirs=config.support_frame_dirs,
|
||||
target_size=target_size,
|
||||
verbose=True)
|
||||
|
||||
max_clips = config.num_videos * config.num_clips_per_video
|
||||
features = extract_features_streaming(video_generator,
|
||||
@@ -220,14 +235,17 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
verbose=True)
|
||||
else:
|
||||
print(f"Extracting features from {len(videos)} video tensors...")
|
||||
features = extractor.extract_features(videos, batch_size=config.batch_size, verbose=True)
|
||||
features = extractor.extract_features(videos,
|
||||
batch_size=config.batch_size,
|
||||
verbose=True)
|
||||
features = features.numpy()
|
||||
|
||||
# Validate feature count
|
||||
expected_count = config.num_videos * config.num_clips_per_video
|
||||
if len(features) < expected_count:
|
||||
raise ValueError(f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
|
||||
f"Found fewer videos than expected. Check your video directory.")
|
||||
raise ValueError(
|
||||
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
|
||||
f"Found fewer videos than expected. Check your video directory.")
|
||||
elif len(features) > expected_count:
|
||||
print(f"Truncating {len(features)} features to {expected_count}")
|
||||
features = features[:expected_count]
|
||||
@@ -294,7 +312,9 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
|
||||
|
||||
# Initialize Extractor using Factory
|
||||
if verbose:
|
||||
print(f"\nInitializing {config.extractor_model.upper()} model on {config.device}...")
|
||||
print(
|
||||
f"\nInitializing {config.extractor_model.upper()} model on {config.device}..."
|
||||
)
|
||||
|
||||
extractor = load_extractor(config.extractor_model, device=config.device)
|
||||
|
||||
@@ -304,11 +324,12 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
|
||||
print("Extracting REAL video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
real_features = load_or_compute_features(videos=real_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=config.cache_real_features,
|
||||
cache_name="real_features")
|
||||
real_features = load_or_compute_features(
|
||||
videos=real_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=config.cache_real_features,
|
||||
cache_name="real_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
|
||||
+28
-10
@@ -19,12 +19,16 @@ class I3DFeatureExtractor(nn.Module):
|
||||
REPO_ID = 'flateon/FVD-I3D-torchscript'
|
||||
MODEL_FILENAME = 'i3d_torchscript.pt'
|
||||
|
||||
def __init__(self, device: str = 'cuda', cache_dir: str | Path | None = None):
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
cache_dir: str | Path | None = None):
|
||||
super().__init__()
|
||||
|
||||
self.device_str = device
|
||||
if device == 'cuda' and not torch.cuda.is_available():
|
||||
print("Warning: CUDA requested but not available – falling back to CPU")
|
||||
print(
|
||||
"Warning: CUDA requested but not available – falling back to CPU"
|
||||
)
|
||||
self.device = torch.device('cpu')
|
||||
else:
|
||||
self.device = torch.device(device)
|
||||
@@ -47,7 +51,9 @@ class I3DFeatureExtractor(nn.Module):
|
||||
|
||||
try:
|
||||
# Download model from Hugging Face Hub
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID, filename=self.MODEL_FILENAME, cache_dir=self.cache_dir)
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
|
||||
# Load directly to chosen device
|
||||
model = torch.jit.load(model_path, map_location=self.device)
|
||||
@@ -55,9 +61,10 @@ class I3DFeatureExtractor(nn.Module):
|
||||
return model
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
raise RuntimeError(
|
||||
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
@@ -81,7 +88,10 @@ class I3DFeatureExtractor(nn.Module):
|
||||
# Resize to 224x224 if needed
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.reshape(B * T, C, H, W)
|
||||
videos = F.interpolate(videos, size=(224, 224), mode='bilinear', align_corners=False)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.reshape(B, T, C, 224, 224)
|
||||
|
||||
# Convert to [B, C, T, H, W] format
|
||||
@@ -90,7 +100,10 @@ class I3DFeatureExtractor(nn.Module):
|
||||
return videos
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self, videos: torch.Tensor, batch_size: int = 32, verbose: bool = True) -> torch.Tensor:
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract I3D features
|
||||
|
||||
@@ -114,11 +127,16 @@ class I3DFeatureExtractor(nn.Module):
|
||||
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
|
||||
|
||||
# Use the HF model without rescale/resize (we handle it in preprocess)
|
||||
features = self.model(batch, rescale=False, resize=False, return_features=True)
|
||||
features = self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
all_features.append(features.cpu())
|
||||
|
||||
return torch.cat(all_features, dim=0)
|
||||
|
||||
def __call__(self, videos: torch.Tensor, batch_size: int = 32) -> torch.Tensor:
|
||||
def __call__(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32) -> torch.Tensor:
|
||||
return self.extract_features(videos, batch_size=batch_size)
|
||||
|
||||
@@ -36,7 +36,10 @@ def main() -> None:
|
||||
cache_real_features=str(script_dir / f'fvd-cache/{model_name}'),
|
||||
)
|
||||
|
||||
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=False)
|
||||
results = compute_fvd_with_config(real_dir,
|
||||
gen_dir,
|
||||
cfg,
|
||||
verbose=False)
|
||||
print(f"FVD: {results['fvd']}\nModel: {results['model']}")
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -59,7 +59,10 @@ def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
print("TEST 1: Identity Test")
|
||||
print("=" * 70)
|
||||
|
||||
result1 = compute_fvd_with_config(real_videos=str(subset_a), gen_videos=str(subset_a), config=config, verbose=False)
|
||||
result1 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_a),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_identity = result1['fvd']
|
||||
print(f"\nIdentity FVD: {fvd_identity:.2f}")
|
||||
|
||||
@@ -67,7 +70,10 @@ def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
print("TEST 2: Real vs Real")
|
||||
print("=" * 70)
|
||||
|
||||
result2 = compute_fvd_with_config(real_videos=str(subset_a), gen_videos=str(subset_b), config=config, verbose=False)
|
||||
result2 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_b),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_real = result2['fvd']
|
||||
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
|
||||
|
||||
@@ -81,7 +87,9 @@ def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
def main() -> None:
|
||||
bair_dir = Path('benchmarks/data/bair_full_videos')
|
||||
|
||||
subset_a, subset_b, count = split_videos(bair_dir, n_per_subset=128, seed=42)
|
||||
subset_a, subset_b, count = split_videos(bair_dir,
|
||||
n_per_subset=128,
|
||||
seed=42)
|
||||
validate_fvd(subset_a, subset_b, count)
|
||||
|
||||
|
||||
|
||||
@@ -54,7 +54,8 @@ def _load_video_cv2(video_path: str | Path,
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1, 2).float() # [T, C, H, W]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
return frames
|
||||
|
||||
if total_frames == 0:
|
||||
@@ -62,11 +63,14 @@ def _load_video_cv2(video_path: str | Path,
|
||||
|
||||
# Determine frame indices for sampling
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
frame_indices = list(range(
|
||||
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames, dtype=int).tolist()
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(np.random.choice(total_frames, num_frames, replace=False))
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
@@ -90,15 +94,17 @@ def _load_video_cv2(video_path: str | Path,
|
||||
cap.release()
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1, 2).float() # [T, C, H, W]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _load_video_from_frames(frame_dir: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform',
|
||||
frame_extensions: list[str] | None = None) -> torch.Tensor:
|
||||
def _load_video_from_frames(
|
||||
frame_dir: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform',
|
||||
frame_extensions: list[str] | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Load video from directory of frame images.
|
||||
|
||||
@@ -125,7 +131,9 @@ def _load_video_from_frames(frame_dir: str | Path,
|
||||
frame_files.extend(frame_dir.glob(f"*{ext}"))
|
||||
|
||||
if len(frame_files) == 0:
|
||||
raise ValueError(f"No frames found in {frame_dir} with extensions {frame_extensions}")
|
||||
raise ValueError(
|
||||
f"No frames found in {frame_dir} with extensions {frame_extensions}"
|
||||
)
|
||||
|
||||
frame_files = sorted(frame_files, key=lambda x: x.name)
|
||||
total_frames = len(frame_files)
|
||||
@@ -135,11 +143,16 @@ def _load_video_from_frames(frame_dir: str | Path,
|
||||
frame_indices = list(range(total_frames))
|
||||
else:
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
|
||||
num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames, dtype=int).tolist()
|
||||
frame_indices = np.linspace(0,
|
||||
total_frames - 1,
|
||||
num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(np.random.choice(total_frames, num_frames, replace=False))
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
@@ -157,7 +170,8 @@ def _load_video_from_frames(frame_dir: str | Path,
|
||||
|
||||
# Stack and convert to tensor
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1, 2).float() # [T, C, H, W]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
@@ -212,12 +226,13 @@ def load_video_auto(video_path: str | Path,
|
||||
raise ValueError(f"Unknown video format at {video_path}")
|
||||
|
||||
|
||||
def sample_clips_from_video(video: torch.Tensor,
|
||||
num_frames_per_clip: int = 16,
|
||||
num_clips: int = 1,
|
||||
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
|
||||
frame_stride: int = 1,
|
||||
temporal_stride: int = 1) -> list[torch.Tensor]:
|
||||
def sample_clips_from_video(
|
||||
video: torch.Tensor,
|
||||
num_frames_per_clip: int = 16,
|
||||
num_clips: int = 1,
|
||||
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
|
||||
frame_stride: int = 1,
|
||||
temporal_stride: int = 1) -> list[torch.Tensor]:
|
||||
"""
|
||||
Sample clips from a video with various strategies.
|
||||
|
||||
@@ -281,7 +296,10 @@ def sample_clips_from_video(video: torch.Tensor,
|
||||
elif strategy == ClipSamplingStrategy.RANDOM:
|
||||
# Sample N random clips
|
||||
for _ in range(num_clips):
|
||||
start = 0 if effective_clip_length == T else np.random.randint(0, T - effective_clip_length + 1)
|
||||
if effective_clip_length == T:
|
||||
start = 0
|
||||
else:
|
||||
start = np.random.randint(0, T - effective_clip_length + 1)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
@@ -294,7 +312,8 @@ def sample_clips_from_video(video: torch.Tensor,
|
||||
clips.append(clip)
|
||||
else:
|
||||
# Multiple uniformly spaced clips
|
||||
step = (T - effective_clip_length) / (num_clips - 1) if num_clips > 1 else 0
|
||||
step = (T - effective_clip_length) / (num_clips -
|
||||
1) if num_clips > 1 else 0
|
||||
for i in range(num_clips):
|
||||
start = int(i * step)
|
||||
start = min(start, T - effective_clip_length)
|
||||
@@ -403,12 +422,15 @@ def load_video_clips_streaming(directory: str | Path,
|
||||
failed_count = 0
|
||||
total_clips = 0
|
||||
|
||||
iterator = tqdm(video_paths, desc="Loading videos") if verbose else video_paths
|
||||
iterator = tqdm(video_paths,
|
||||
desc="Loading videos") if verbose else video_paths
|
||||
|
||||
for video_path in iterator:
|
||||
try:
|
||||
# Load full video
|
||||
video = load_video_auto(video_path, num_frames=None, sample_strategy='uniform')
|
||||
video = load_video_auto(video_path,
|
||||
num_frames=None,
|
||||
sample_strategy='uniform')
|
||||
|
||||
# Sample clips from video
|
||||
clips = sample_clips_from_video(video,
|
||||
@@ -423,14 +445,18 @@ def load_video_clips_streaming(directory: str | Path,
|
||||
T, C, H, W = clip.shape
|
||||
if target_size != (H, W):
|
||||
# Resize to target size
|
||||
clip = clip.contiguous() # Fix non-contiguous tensors first
|
||||
clip_flat = clip.view(T * C, H, W).unsqueeze(0) # [1, T*C, H, W]
|
||||
clip_resized = torch.nn.functional.interpolate(clip_flat,
|
||||
size=target_size,
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
clip = clip_resized.squeeze(0).view(T, C, target_size[0],
|
||||
target_size[1]) # Back to [T, C, H, W]
|
||||
clip = clip.contiguous(
|
||||
) # Fix non-contiguous tensors first
|
||||
clip_flat = clip.view(T * C, H,
|
||||
W).unsqueeze(0) # [1, T*C, H, W]
|
||||
clip_resized = torch.nn.functional.interpolate(
|
||||
clip_flat,
|
||||
size=target_size,
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
clip = clip_resized.squeeze(0).view(
|
||||
T, C, target_size[0],
|
||||
target_size[1]) # Back to [T, C, H, W]
|
||||
resized_clips.append(clip)
|
||||
clips = resized_clips
|
||||
|
||||
@@ -454,7 +480,11 @@ def load_video_clips_streaming(directory: str | Path,
|
||||
|
||||
failure_rate = failed_count / len(video_paths)
|
||||
if failure_rate > 0.1: # More than 10% failed
|
||||
print(f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})")
|
||||
print(
|
||||
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
|
||||
)
|
||||
|
||||
if verbose:
|
||||
print(f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos")
|
||||
print(
|
||||
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
|
||||
)
|
||||
|
||||
+63
-30
@@ -100,7 +100,10 @@ DEFAULT_PIP_PATTERNS = {
|
||||
def run(command):
|
||||
"""Return (return-code, stdout, stderr)."""
|
||||
shell = True if type(command) is str else False
|
||||
p = subprocess.Popen(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=shell)
|
||||
p = subprocess.Popen(command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
shell=shell)
|
||||
raw_output, raw_err = p.communicate()
|
||||
rc = p.returncode
|
||||
if get_platform() == 'win32':
|
||||
@@ -154,7 +157,8 @@ def get_conda_packages(run_lambda, patterns=None):
|
||||
return out
|
||||
|
||||
return "\n".join(line for line in out.splitlines()
|
||||
if not line.startswith("#") and any(name in line for name in patterns))
|
||||
if not line.startswith("#") and any(name in line
|
||||
for name in patterns))
|
||||
|
||||
|
||||
def get_gcc_version(run_lambda):
|
||||
@@ -162,24 +166,27 @@ def get_gcc_version(run_lambda):
|
||||
|
||||
|
||||
def get_clang_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'clang --version', r'clang version (.*)')
|
||||
return run_and_parse_first_match(run_lambda, 'clang --version',
|
||||
r'clang version (.*)')
|
||||
|
||||
|
||||
def get_cmake_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'cmake --version', r'cmake (.*)')
|
||||
return run_and_parse_first_match(run_lambda, 'cmake --version',
|
||||
r'cmake (.*)')
|
||||
|
||||
|
||||
def get_nvidia_driver_version(run_lambda):
|
||||
if get_platform() == 'darwin':
|
||||
cmd = 'kextstat | grep -i cuda'
|
||||
return run_and_parse_first_match(run_lambda, cmd, r'com[.]nvidia[.]CUDA [(](.*?)[)]')
|
||||
return run_and_parse_first_match(run_lambda, cmd,
|
||||
r'com[.]nvidia[.]CUDA [(](.*?)[)]')
|
||||
smi = get_nvidia_smi()
|
||||
return run_and_parse_first_match(run_lambda, smi, r'Driver Version: (.*?) ')
|
||||
|
||||
|
||||
def get_gpu_info(run_lambda):
|
||||
if get_platform() == 'darwin' or (TORCH_AVAILABLE and hasattr(torch.version, 'hip')
|
||||
and torch.version.hip is not None):
|
||||
if get_platform() == 'darwin' or (TORCH_AVAILABLE and hasattr(
|
||||
torch.version, 'hip') and torch.version.hip is not None):
|
||||
if TORCH_AVAILABLE and torch.cuda.is_available():
|
||||
if torch.version.hip is not None:
|
||||
prop = torch.cuda.get_device_properties(0)
|
||||
@@ -201,7 +208,8 @@ def get_gpu_info(run_lambda):
|
||||
|
||||
|
||||
def get_running_cuda_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'nvcc --version', r'release .+ V(.*)')
|
||||
return run_and_parse_first_match(run_lambda, 'nvcc --version',
|
||||
r'release .+ V(.*)')
|
||||
|
||||
|
||||
def get_cudnn_version(run_lambda):
|
||||
@@ -247,7 +255,8 @@ def get_nvidia_smi():
|
||||
if get_platform() == 'win32':
|
||||
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
|
||||
program_files_root = os.environ.get('PROGRAMFILES', 'C:\\Program Files')
|
||||
legacy_path = os.path.join(program_files_root, 'NVIDIA Corporation', 'NVSMI', smi)
|
||||
legacy_path = os.path.join(program_files_root, 'NVIDIA Corporation',
|
||||
'NVSMI', smi)
|
||||
new_path = os.path.join(system_root, 'System32', smi)
|
||||
smis = [new_path, legacy_path]
|
||||
for candidate_smi in smis:
|
||||
@@ -379,8 +388,10 @@ def get_cpu_info(run_lambda):
|
||||
if get_platform() == 'linux':
|
||||
rc, out, err = run_lambda('lscpu')
|
||||
elif get_platform() == 'win32':
|
||||
rc, out, err = run_lambda('wmic cpu get Name,Manufacturer,Family,Architecture,ProcessorType,DeviceID, \
|
||||
CurrentClockSpeed,MaxClockSpeed,L2CacheSize,L2CacheSpeed,Revision /VALUE')
|
||||
rc, out, err = run_lambda(
|
||||
'wmic cpu get Name,Manufacturer,Family,Architecture,ProcessorType,DeviceID, \
|
||||
CurrentClockSpeed,MaxClockSpeed,L2CacheSize,L2CacheSpeed,Revision /VALUE'
|
||||
)
|
||||
elif get_platform() == 'darwin':
|
||||
rc, out, err = run_lambda("sysctl -n machdep.cpu.brand_string")
|
||||
cpu_info = 'None'
|
||||
@@ -405,22 +416,27 @@ def get_platform():
|
||||
|
||||
|
||||
def get_mac_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'sw_vers -productVersion', r'(.*)')
|
||||
return run_and_parse_first_match(run_lambda, 'sw_vers -productVersion',
|
||||
r'(.*)')
|
||||
|
||||
|
||||
def get_windows_version(run_lambda):
|
||||
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
|
||||
wmic_cmd = os.path.join(system_root, 'System32', 'Wbem', 'wmic')
|
||||
findstr_cmd = os.path.join(system_root, 'System32', 'findstr')
|
||||
return run_and_read_all(run_lambda, '{} os get Caption | {} /v Caption'.format(wmic_cmd, findstr_cmd))
|
||||
return run_and_read_all(
|
||||
run_lambda,
|
||||
'{} os get Caption | {} /v Caption'.format(wmic_cmd, findstr_cmd))
|
||||
|
||||
|
||||
def get_lsb_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'lsb_release -a', r'Description:\t(.*)')
|
||||
return run_and_parse_first_match(run_lambda, 'lsb_release -a',
|
||||
r'Description:\t(.*)')
|
||||
|
||||
|
||||
def check_release_file(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'cat /etc/*-release', r'PRETTY_NAME="(.*)"')
|
||||
return run_and_parse_first_match(run_lambda, 'cat /etc/*-release',
|
||||
r'PRETTY_NAME="(.*)"')
|
||||
|
||||
|
||||
def get_os(run_lambda):
|
||||
@@ -483,10 +499,13 @@ def get_pip_packages(run_lambda, patterns=None):
|
||||
elif shutil.which("uv") is not None:
|
||||
cmd = ["uv", "pip", "list", "--format=freeze"]
|
||||
else:
|
||||
raise RuntimeError("Could not collect pip list output (pip or uv module not available)")
|
||||
raise RuntimeError(
|
||||
"Could not collect pip list output (pip or uv module not available)"
|
||||
)
|
||||
|
||||
out = run_and_read_all(run_lambda, cmd)
|
||||
return "\n".join(line for line in out.splitlines() if any(name in line for name in patterns))
|
||||
return "\n".join(line for line in out.splitlines()
|
||||
if any(name in line for name in patterns))
|
||||
|
||||
pip_version = 'pip3' if sys.version[0] == '3' else 'pip'
|
||||
out = run_with_pip()
|
||||
@@ -518,7 +537,8 @@ def is_xnnpack_available():
|
||||
def get_env_vars():
|
||||
env_vars = ''
|
||||
secret_terms = ('secret', 'token', 'api', 'access', 'password')
|
||||
report_prefix = ("TORCH", "NCCL", "PYTORCH", "CUDA", "CUBLAS", "CUDNN", "OMP_", "MKL_", "NVIDIA")
|
||||
report_prefix = ("TORCH", "NCCL", "PYTORCH", "CUDA", "CUBLAS", "CUDNN",
|
||||
"OMP_", "MKL_", "NVIDIA")
|
||||
for k, v in os.environ.items():
|
||||
if any(term in k.lower() for term in secret_terms):
|
||||
continue
|
||||
@@ -539,7 +559,8 @@ def get_env_info():
|
||||
debug_mode_str = str(torch.version.debug)
|
||||
cuda_available_str = str(torch.cuda.is_available())
|
||||
cuda_version_str = torch.version.cuda
|
||||
if not hasattr(torch.version, 'hip') or torch.version.hip is None: # cuda version
|
||||
if not hasattr(torch.version,
|
||||
'hip') or torch.version.hip is None: # cuda version
|
||||
hip_compiled_version = hip_runtime_version = miopen_runtime_version = 'N/A'
|
||||
else: # HIP version
|
||||
|
||||
@@ -567,8 +588,9 @@ def get_env_info():
|
||||
return SystemEnv(
|
||||
torch_version=version_str,
|
||||
is_debug_build=debug_mode_str,
|
||||
python_version='{} ({}-bit runtime)'.format(sys_version,
|
||||
sys.maxsize.bit_length() + 1),
|
||||
python_version='{} ({}-bit runtime)'.format(
|
||||
sys_version,
|
||||
sys.maxsize.bit_length() + 1),
|
||||
python_platform=get_python_platform(),
|
||||
is_cuda_available=cuda_available_str,
|
||||
cuda_compiled_version=cuda_version_str,
|
||||
@@ -692,8 +714,10 @@ def pretty_str(envinfo):
|
||||
'nvidia_driver_version',
|
||||
]
|
||||
all_cuda_fields = dynamic_cuda_fields + ['cudnn_version']
|
||||
all_dynamic_cuda_fields_missing = all(mutable_dict[field] is None for field in dynamic_cuda_fields)
|
||||
if TORCH_AVAILABLE and not torch.cuda.is_available() and all_dynamic_cuda_fields_missing:
|
||||
all_dynamic_cuda_fields_missing = all(mutable_dict[field] is None
|
||||
for field in dynamic_cuda_fields)
|
||||
if TORCH_AVAILABLE and not torch.cuda.is_available(
|
||||
) and all_dynamic_cuda_fields_missing:
|
||||
for field in all_cuda_fields:
|
||||
mutable_dict[field] = 'No CUDA'
|
||||
if envinfo.cuda_compiled_version is None:
|
||||
@@ -706,15 +730,19 @@ def pretty_str(envinfo):
|
||||
mutable_dict = replace_nones(mutable_dict)
|
||||
|
||||
# If either of these are '', replace with 'No relevant packages'
|
||||
mutable_dict['pip_packages'] = replace_if_empty(mutable_dict['pip_packages'])
|
||||
mutable_dict['conda_packages'] = replace_if_empty(mutable_dict['conda_packages'])
|
||||
mutable_dict['pip_packages'] = replace_if_empty(
|
||||
mutable_dict['pip_packages'])
|
||||
mutable_dict['conda_packages'] = replace_if_empty(
|
||||
mutable_dict['conda_packages'])
|
||||
|
||||
# Tag conda and pip packages with a prefix
|
||||
# If they were previously None, they'll show up as ie '[conda] Could not collect'
|
||||
if mutable_dict['pip_packages']:
|
||||
mutable_dict['pip_packages'] = prepend(mutable_dict['pip_packages'], '[{}] '.format(envinfo.pip_version))
|
||||
mutable_dict['pip_packages'] = prepend(
|
||||
mutable_dict['pip_packages'], '[{}] '.format(envinfo.pip_version))
|
||||
if mutable_dict['conda_packages']:
|
||||
mutable_dict['conda_packages'] = prepend(mutable_dict['conda_packages'], '[conda] ')
|
||||
mutable_dict['conda_packages'] = prepend(mutable_dict['conda_packages'],
|
||||
'[conda] ')
|
||||
mutable_dict['cpu_info'] = envinfo.cpu_info
|
||||
return env_info_fmt.format(**mutable_dict)
|
||||
|
||||
@@ -728,13 +756,18 @@ def main():
|
||||
output = get_pretty_env_info()
|
||||
print(output)
|
||||
|
||||
if TORCH_AVAILABLE and hasattr(torch, 'utils') and hasattr(torch.utils, '_crash_handler'):
|
||||
if TORCH_AVAILABLE and hasattr(torch, 'utils') and hasattr(
|
||||
torch.utils, '_crash_handler'):
|
||||
minidump_dir = torch.utils._crash_handler.DEFAULT_MINIDUMP_DIR
|
||||
if sys.platform == "linux" and os.path.exists(minidump_dir):
|
||||
dumps = [os.path.join(minidump_dir, dump) for dump in os.listdir(minidump_dir)]
|
||||
dumps = [
|
||||
os.path.join(minidump_dir, dump)
|
||||
for dump in os.listdir(minidump_dir)
|
||||
]
|
||||
latest = max(dumps, key=os.path.getctime)
|
||||
ctime = os.path.getctime(latest)
|
||||
creation_time = datetime.datetime.fromtimestamp(ctime).strftime('%Y-%m-%d %H:%M:%S')
|
||||
creation_time = datetime.datetime.fromtimestamp(ctime).strftime(
|
||||
'%Y-%m-%d %H:%M:%S')
|
||||
msg = "\n*** Detected a minidump at {} created on {}, ".format(latest, creation_time) + \
|
||||
"if this is related to your bug please include it when you file a report ***"
|
||||
print(msg, file=sys.stderr)
|
||||
|
||||
+3
-2
@@ -1,4 +1,5 @@
|
||||
from .video_generator.nodes import (NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS)
|
||||
from .video_generator.nodes import (NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||||
@@ -14,7 +14,10 @@ class LoadImagePath:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
|
||||
files = [
|
||||
f for f in os.listdir(input_dir)
|
||||
if os.path.isfile(os.path.join(input_dir, f))
|
||||
]
|
||||
files = folder_paths.filter_files_content_types(files, ["image"])
|
||||
return {
|
||||
"required": {
|
||||
@@ -62,10 +65,13 @@ class LoadImagePath:
|
||||
None,
|
||||
]
|
||||
if 'A' in processed_image.getbands():
|
||||
mask = np.array(processed_image.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = np.array(processed_image.getchannel('A')).astype(
|
||||
np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
elif processed_image.mode == 'P' and 'transparency' in processed_image.info:
|
||||
mask = np.array(processed_image.convert('RGBA').getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = np.array(
|
||||
processed_image.convert('RGBA').getchannel('A')).astype(
|
||||
np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
|
||||
@@ -9,7 +9,8 @@ from PIL import ImageFile, UnidentifiedImageError
|
||||
T = TypeVar('T')
|
||||
|
||||
|
||||
def conditioning_set_values(conditioning: list[Any], values: dict[str, Any] | None = None) -> list[Any]:
|
||||
def conditioning_set_values(conditioning: list[Any],
|
||||
values: dict[str, Any] | None = None) -> list[Any]:
|
||||
if values is None:
|
||||
values = {}
|
||||
c = []
|
||||
@@ -26,7 +27,8 @@ def pillow(fn: Callable[[Any], T], arg: Any) -> T:
|
||||
prev_value = None
|
||||
try:
|
||||
x = fn(arg)
|
||||
except (OSError, UnidentifiedImageError, ValueError): #PIL issues #4472 and #2445, also fixes ComfyUI issue #3416
|
||||
except (OSError, UnidentifiedImageError, ValueError
|
||||
): #PIL issues #4472 and #2445, also fixes ComfyUI issue #3416
|
||||
prev_value = ImageFile.LOAD_TRUNCATED_IMAGES
|
||||
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||
x = fn(arg)
|
||||
@@ -37,7 +39,12 @@ def pillow(fn: Callable[[Any], T], arg: Any) -> T:
|
||||
|
||||
|
||||
def hasher() -> Callable[[], Any]:
|
||||
hashfuncs = {"md5": hashlib.md5, "sha1": hashlib.sha1, "sha256": hashlib.sha256, "sha512": hashlib.sha512}
|
||||
hashfuncs = {
|
||||
"md5": hashlib.md5,
|
||||
"sha1": hashlib.sha1,
|
||||
"sha256": hashlib.sha256,
|
||||
"sha512": hashlib.sha512
|
||||
}
|
||||
return hashfuncs[args.default_hashing_function]
|
||||
|
||||
|
||||
@@ -51,7 +58,8 @@ def string_to_torch_dtype(string: str) -> torch.dtype | None:
|
||||
return None
|
||||
|
||||
|
||||
def image_alpha_fix(destination: torch.Tensor, source: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def image_alpha_fix(destination: torch.Tensor,
|
||||
source: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if destination.shape[-1] < source.shape[-1]:
|
||||
source = source[..., :destination.shape[-1]]
|
||||
elif destination.shape[-1] > source.shape[-1]:
|
||||
|
||||
@@ -26,7 +26,11 @@ class TextEncoderConfig:
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
def set_args(self, prefix, quant_config, lora_config):
|
||||
raw_args = {"prefix": prefix, "quant_config": quant_config, "lora_config": lora_config}
|
||||
raw_args = {
|
||||
"prefix": prefix,
|
||||
"quant_config": quant_config,
|
||||
"lora_config": lora_config
|
||||
}
|
||||
|
||||
# Filter out keys where value is -99999
|
||||
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
|
||||
|
||||
@@ -13,7 +13,11 @@ from comfy.model_management import processing_interrupted
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo import VideoGenerator as FastVideoGenerator
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))))
|
||||
sys.path.insert(
|
||||
0,
|
||||
os.path.dirname(
|
||||
os.path.dirname(
|
||||
os.path.dirname(os.path.dirname(os.path.abspath(__file__))))))
|
||||
|
||||
|
||||
# Custom exception for interruption
|
||||
@@ -24,7 +28,8 @@ class GenerationInterruptedException(Exception):
|
||||
# Custom exception for interruption that ComfyUI will recognize
|
||||
class GenerationCancelledException(Exception):
|
||||
|
||||
def __init__(self, message: str = "Generation was cancelled by user") -> None:
|
||||
def __init__(self,
|
||||
message: str = "Generation was cancelled by user") -> None:
|
||||
self.message = message
|
||||
super().__init__(self.message)
|
||||
|
||||
@@ -134,7 +139,8 @@ class VideoGenerator:
|
||||
self._generation_interrupted = True
|
||||
|
||||
# Try to send interrupt signal to worker processes
|
||||
if self.generator is not None and hasattr(self.generator, 'executor'):
|
||||
if self.generator is not None and hasattr(
|
||||
self.generator, 'executor'):
|
||||
try:
|
||||
# The MultiprocExecutor has a workers attribute
|
||||
if hasattr(self.generator.executor, 'workers'):
|
||||
@@ -150,12 +156,16 @@ class VideoGenerator:
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
def _run_generation(self, prompt: str, output_path: str, inference_args: dict[str, Any]) -> None:
|
||||
def _run_generation(self, prompt: str, output_path: str,
|
||||
inference_args: dict[str, Any]) -> None:
|
||||
"""Thread function to run the generation"""
|
||||
try:
|
||||
if self.generator is not None:
|
||||
self.generator.generate_video(prompt=prompt, output_path=output_path, **inference_args)
|
||||
self._generation_result = os.path.join(output_path, f"{prompt[:100]}.mp4")
|
||||
self.generator.generate_video(prompt=prompt,
|
||||
output_path=output_path,
|
||||
**inference_args)
|
||||
self._generation_result = os.path.join(output_path,
|
||||
f"{prompt[:100]}.mp4")
|
||||
else:
|
||||
raise RuntimeError("Generator is not initialized")
|
||||
except Exception as e:
|
||||
@@ -216,7 +226,8 @@ class VideoGenerator:
|
||||
update_config_from_args(pipeline_config.vae_config, vae_config)
|
||||
|
||||
if text_encoder_config is not None:
|
||||
update_config_from_args(pipeline_config.text_encoder_configs, text_encoder_config)
|
||||
update_config_from_args(pipeline_config.text_encoder_configs,
|
||||
text_encoder_config)
|
||||
|
||||
# Update top-level pipeline config with remaining arguments
|
||||
raw_pipeline_args = {}
|
||||
@@ -234,7 +245,10 @@ class VideoGenerator:
|
||||
raw_pipeline_args['text_encoder_precision'] = text_encoder_precision
|
||||
|
||||
# Filter out any value explicitly set to -99999 (auto values)
|
||||
pipeline_args = {k: v for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)}
|
||||
pipeline_args = {
|
||||
k: v
|
||||
for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
update_config_from_args(pipeline_config, pipeline_args)
|
||||
|
||||
@@ -248,30 +262,38 @@ class VideoGenerator:
|
||||
if dit_cpu_offload is not None:
|
||||
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
|
||||
|
||||
generation_args = {k: v for k, v in raw_generation_args.items() if str(int(v)) != str(-99999)}
|
||||
generation_args = {
|
||||
k: v
|
||||
for k, v in raw_generation_args.items()
|
||||
if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
if self.generator is None:
|
||||
print('generation_args', generation_args)
|
||||
print('pipeline_config', pipeline_config)
|
||||
self.generator = FastVideoGenerator.from_pretrained(model_path=model_path,
|
||||
**generation_args,
|
||||
pipeline_config=pipeline_config)
|
||||
self.generator = FastVideoGenerator.from_pretrained(
|
||||
model_path=model_path,
|
||||
**generation_args,
|
||||
pipeline_config=pipeline_config)
|
||||
|
||||
print('inference_args', inference_args)
|
||||
|
||||
# Start a thread to run the generation
|
||||
self._generation_thread = threading.Thread(target=self._run_generation,
|
||||
args=(prompt, output_path, inference_args),
|
||||
args=(prompt, output_path,
|
||||
inference_args),
|
||||
daemon=True)
|
||||
self._generation_thread.start()
|
||||
|
||||
# Start a background thread to monitor for interruptions
|
||||
self._generation_active = True
|
||||
self._interrupt_thread = threading.Thread(target=self._monitor_for_interruption, daemon=True)
|
||||
self._interrupt_thread = threading.Thread(
|
||||
target=self._monitor_for_interruption, daemon=True)
|
||||
self._interrupt_thread.start()
|
||||
|
||||
# Wait for either completion or interruption
|
||||
while self._generation_thread.is_alive() and not self._interrupt_event.is_set():
|
||||
while self._generation_thread.is_alive(
|
||||
) and not self._interrupt_event.is_set():
|
||||
self._generation_thread.join(timeout=0.5)
|
||||
|
||||
self._generation_active = False
|
||||
|
||||
+1
-1
@@ -36,4 +36,4 @@ Then open your browser to: http://localhost:8000
|
||||
|
||||
## Automatic Deployment
|
||||
|
||||
Documentation is automatically built and deployed to GitHub Pages when changes are pushed to the `main` branch via the `.github/workflows/infra-docs.yml` workflow.
|
||||
Documentation is automatically built and deployed to GitHub Pages when changes are pushed to the `main` branch via the `.github/workflows/docs.yml` workflow.
|
||||
|
||||
@@ -1,306 +0,0 @@
|
||||
# CI Architecture
|
||||
|
||||
## Overview
|
||||
|
||||
FastVideo uses a three-tier CI pipeline designed to keep feedback fast on every push while
|
||||
protecting `main` through a full GPU regression suite before any merge.
|
||||
|
||||
```
|
||||
PR push
|
||||
│
|
||||
├─► Tier 1: Pre-commit (~2 min)
|
||||
│ GitHub Actions / ubuntu-latest
|
||||
│ yapf, ruff, mypy, codespell, pymarkdown, actionlint, check-filenames
|
||||
│
|
||||
└─► Tier 2: Fastcheck (~10-20 min, path-filtered)
|
||||
Buildkite / Modal GPU instances
|
||||
Only runs tests for paths you changed
|
||||
|
||||
│ (developer comments /merge or maintainer adds 'ready' label)
|
||||
▼
|
||||
Tier 3: Full Suite (~60-90 min)
|
||||
Buildkite / Modal GPU instances
|
||||
All integration, SSIM, training, and performance tests
|
||||
Runs on the PR branch directly
|
||||
│
|
||||
pass ──► Mergify auto-squash-merges to main, branch deleted
|
||||
fail ──► Mergify removes 'ready' label; fix and /merge again
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## CI Tiers
|
||||
|
||||
### Tier 1: Pre-commit (every PR push)
|
||||
|
||||
| Attribute | Value |
|
||||
|-----------|-------|
|
||||
| Triggered by | Every push to any PR branch, plus pushes to `main` |
|
||||
| Runs on | GitHub Actions, `ubuntu-latest` |
|
||||
| Duration | ~2 minutes |
|
||||
|
||||
**Checks run** (from `.pre-commit-config.yaml`, stage: `manual`):
|
||||
|
||||
| Hook | What it checks |
|
||||
|------|---------------|
|
||||
| `yapf` | Python code formatting |
|
||||
| `ruff` | Python linting and auto-fixable style issues |
|
||||
| `codespell` | Spelling errors in code and docs |
|
||||
| `pymarkdown` | Markdown formatting |
|
||||
| `actionlint` | GitHub Actions workflow syntax |
|
||||
| `mypy` | Static type checking (Python 3.10 target) |
|
||||
| `check-filenames` | No spaces in tracked filenames |
|
||||
|
||||
A failure here means code style or type issues. Run `pre-commit run --all-files` locally to
|
||||
replicate CI results before pushing.
|
||||
|
||||
---
|
||||
|
||||
### Tier 2: Fastcheck (path-filtered, every PR push)
|
||||
|
||||
| Attribute | Value |
|
||||
|-----------|-------|
|
||||
| Triggered by | Every push; the monorepo-diff plugin skips tests for unchanged paths |
|
||||
| Runs on | Buildkite, Modal GPU instances |
|
||||
| Duration | ~10-20 minutes per test (run in parallel) |
|
||||
|
||||
**Tests and their path triggers:**
|
||||
|
||||
| Buildkite label | `TEST_TYPE` | Triggers when you change |
|
||||
|-----------------|-------------|--------------------------|
|
||||
| Encoder Tests | `encoder` | `fastvideo/models/encoders/**`, `fastvideo/models/loader/**`, `fastvideo/tests/encoders/**`, `pyproject.toml`, `docker/Dockerfile.python3.12` |
|
||||
| VAE Tests | `vae` | `fastvideo/models/vaes/**`, `fastvideo/models/loader/**`, `fastvideo/tests/vaes/**`, `pyproject.toml`, `docker/Dockerfile.python3.12` |
|
||||
| Transformer Tests | `transformer` | `fastvideo/models/dits/**`, `fastvideo/models/loader/**`, `fastvideo/tests/transformers/**`, `fastvideo/layers/**`, `fastvideo/attention/**`, `pyproject.toml`, `docker/Dockerfile.python3.12` |
|
||||
| Kernel Tests | `kernel_tests` | `fastvideo-kernel/**`, `pyproject.toml`, `docker/Dockerfile.python3.12` |
|
||||
| Unit Tests | `unit_test` | `fastvideo/**`, `.buildkite/**`, `.github/**`, `pyproject.toml`, `docker/Dockerfile.python3.12` |
|
||||
|
||||
A Fastcheck failure means a component-level regression. Check the Buildkite build log for the
|
||||
failing test's output.
|
||||
|
||||
---
|
||||
|
||||
### Tier 3: Full Test Suite (triggered by `ready` label)
|
||||
|
||||
| Attribute | Value |
|
||||
|-----------|-------|
|
||||
| Triggered by | Adding the `ready` label to the PR (via `/merge` command), or a `/test full` command |
|
||||
| Runs on | Buildkite, Modal GPU instances |
|
||||
| Duration | 60-90 minutes total (tests run in parallel, path-filtered) |
|
||||
|
||||
**Tests:**
|
||||
|
||||
| Buildkite label | `TEST_TYPE` | Timeout |
|
||||
|-----------------|-------------|---------|
|
||||
| SSIM Tests | `ssim` | 90 min |
|
||||
| LoRA Inference Tests | `inference_lora` | 20 min |
|
||||
| Training Tests | `training` | 15 min |
|
||||
| Distillation DMD Tests | `distillation_dmd` | 15 min |
|
||||
| Self-Forcing Tests | `self_forcing` | 15 min |
|
||||
| LoRA Training Tests | `training_lora` | 15 min |
|
||||
| Training Tests VSA | `training_vsa` | 15 min |
|
||||
| Inference Tests VMoBA | `inference_vmoba` | 15 min |
|
||||
| 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.
|
||||
|
||||
---
|
||||
|
||||
## Auto-merge Flow
|
||||
|
||||
Mergify prevents untested code from landing on `main` by gating squash-merge on the Full
|
||||
Suite passing directly on the PR branch.
|
||||
|
||||
**How it works:**
|
||||
|
||||
1. A developer comments `/merge` on an approved PR (or a maintainer adds the `ready` label).
|
||||
2. The `ready` label triggers `ci-trigger-full-suite.yml`, which calls the Buildkite API to
|
||||
run the Full Suite on the PR branch itself.
|
||||
3. While the Full Suite runs, Mergify also auto-rebases the PR branch against `main` if it
|
||||
is behind and has no conflicts.
|
||||
4. Once the Full Suite posts `full-suite-passed`, Mergify checks all **merge conditions**:
|
||||
- `pre-commit` check is green
|
||||
- `fastcheck-passed` check is green
|
||||
- `full-suite-passed` check is green
|
||||
- At least 1 approved review (`#approved-reviews-by>=1`)
|
||||
- PR title starts with a valid `[type]` tag
|
||||
- PR is not a draft
|
||||
- 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.
|
||||
|
||||
**Merge conditions summary:**
|
||||
|
||||
| Condition | Meaning |
|
||||
|-----------|---------|
|
||||
| `check-success~=pre-commit` | Tier 1 pre-commit must be green |
|
||||
| `check-success=fastcheck-passed` | Tier 2 Fastcheck must be green |
|
||||
| `check-success=full-suite-passed` | Tier 3 Full Suite must be green |
|
||||
| `#approved-reviews-by>=1` | At least one approved review |
|
||||
| `title~=(?i)^\[(feat|bugfix|...)` | PR title has a valid type tag |
|
||||
| `-draft` | PR is not in draft state |
|
||||
| `-conflict` | No merge conflicts with base branch |
|
||||
| `-closed` | PR is still open |
|
||||
|
||||
---
|
||||
|
||||
## Label System
|
||||
|
||||
Labels are applied automatically. You don't need to set them manually.
|
||||
|
||||
### Type Labels (from PR title prefix)
|
||||
|
||||
Applied by Mergify based on the `[tag]` at the start of the PR title.
|
||||
|
||||
| Label | Matched title prefix | Meaning |
|
||||
|-------|---------------------|---------|
|
||||
| `type: feat` | `[feat]` or `[feature]` | New feature or capability |
|
||||
| `type: bugfix` | `[bugfix]` or `[fix]` | Bug fix |
|
||||
| `type: refactor` | `[refactor]` | Code restructuring, no behavior change |
|
||||
| `type: perf` | `[perf]` | Performance improvement |
|
||||
| `type: ci` | `[ci]` | CI/CD or tooling changes |
|
||||
| `type: docs` | `[doc]` or `[docs]` | Documentation only |
|
||||
| `type: misc` | `[misc]` or `[chore]` | Housekeeping, dependency bumps |
|
||||
| `type: new-model` | `[new-model]` | Adding a new model |
|
||||
|
||||
### Scope Labels (from changed files)
|
||||
|
||||
Applied by Mergify based on which paths you modified. Multiple scope labels can be added.
|
||||
|
||||
| Label | File paths that trigger it |
|
||||
|-------|---------------------------|
|
||||
| `scope: training` | `fastvideo/train/`, `fastvideo/training/`, `fastvideo/distillation/`, `examples/train/`, `examples/training/`, `examples/distill/` |
|
||||
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/configs/sample/`, `fastvideo/configs/pipelines/`, `examples/inference/` |
|
||||
| `scope: attention` | `fastvideo/attention/` |
|
||||
| `scope: kernel` | `fastvideo-kernel/`, `csrc/` |
|
||||
| `scope: data` | `fastvideo/dataset/`, `fastvideo/pipelines/preprocess/`, `examples/preprocessing/` |
|
||||
| `scope: infra` | `.github/`, `.buildkite/`, `fastvideo/tests/`, `docker/` |
|
||||
| `scope: distributed` | `fastvideo/distributed/` |
|
||||
| `scope: docs` | `docs/` |
|
||||
| `scope: ui` | `ui/` |
|
||||
| `scope: model` | `fastvideo/models/`, `fastvideo/layers/`, `fastvideo/configs/models/` |
|
||||
|
||||
### Process Labels
|
||||
|
||||
| Label | Who sets it | Meaning |
|
||||
|-------|-------------|---------|
|
||||
| `ready` | Developer (`/merge` command) or maintainer | Triggers Full Suite and enables auto-merge |
|
||||
| `needs-rebase` | Mergify (automatic) | PR has merge conflicts; rebase needed |
|
||||
| `do-not-merge` | Maintainer | Blocks queue entry regardless of other conditions |
|
||||
|
||||
---
|
||||
|
||||
## PR Title Format
|
||||
|
||||
All PR titles targeting `main` must start with a bracketed type tag. This is enforced by a
|
||||
Mergify merge protection rule and is required before a PR can be squash-merged.
|
||||
|
||||
**Format:**
|
||||
|
||||
```
|
||||
[type] Short description
|
||||
```
|
||||
|
||||
**Valid type tags:**
|
||||
|
||||
`feat`, `feature`, `bugfix`, `fix`, `refactor`, `perf`, `ci`, `doc`, `docs`, `misc`, `chore`,
|
||||
`kernel`, `new-model`
|
||||
|
||||
**Valid examples:**
|
||||
|
||||
```
|
||||
[feat] Add causal Wan pipeline
|
||||
[bugfix] Fix VAE temporal tiling corruption
|
||||
[refactor] Restructure training framework
|
||||
[perf] Optimize FlashAttention kernel dispatch
|
||||
[docs] Add inference guide for LoRA
|
||||
[new-model] Port HunyuanVideo 1.5 to FastVideo
|
||||
```
|
||||
|
||||
**Invalid examples (will block merge):**
|
||||
|
||||
```
|
||||
Add causal Wan pipeline ← missing type tag
|
||||
FEAT: Add pipeline ← wrong format
|
||||
feat: Add pipeline ← square brackets required
|
||||
```
|
||||
|
||||
If your title is invalid, Mergify posts a comment explaining the required format and the
|
||||
merge protection check will remain failed until you update the title.
|
||||
|
||||
---
|
||||
|
||||
## Slash Commands
|
||||
|
||||
Slash commands let contributors and maintainers trigger CI actions directly from PR comments.
|
||||
**Write permission to the repository is required.**
|
||||
|
||||
The command is recognized within a few seconds. The workflow reacts with a 🚀 emoji to confirm.
|
||||
|
||||
### `/merge`
|
||||
|
||||
```
|
||||
/merge
|
||||
```
|
||||
|
||||
Adds the `ready` label to the PR, which triggers the Full Suite on your PR branch and
|
||||
enables Mergify to auto-squash-merge once all conditions pass.
|
||||
|
||||
The command first removes the `ready` label if it is already present, then re-adds it. This
|
||||
ensures the `labeled` event fires and a fresh Full Suite build is started even on a re-try.
|
||||
|
||||
### `/test <name>`
|
||||
|
||||
```
|
||||
/test <name>
|
||||
```
|
||||
|
||||
Triggers a specific Buildkite test or suite on the current PR branch.
|
||||
|
||||
| Command | Runs | Maps to `TEST_TYPE` |
|
||||
|---------|------|---------------------|
|
||||
| `/test encoder` | Encoder Tests (Fastcheck) | `encoder` |
|
||||
| `/test vae` | VAE Tests (Fastcheck) | `vae` |
|
||||
| `/test transformer` | Transformer Tests (Fastcheck) | `transformer` |
|
||||
| `/test kernel` | Kernel Tests (Fastcheck) | `kernel_tests` |
|
||||
| `/test unit` | Unit Tests (Fastcheck) | `unit_test` |
|
||||
| `/test ssim` | SSIM regression tests | `ssim` |
|
||||
| `/test training` | Training pipeline tests | `training` |
|
||||
| `/test lora-inference` | LoRA inference tests | `inference_lora` |
|
||||
| `/test lora-training` | LoRA training tests | `training_lora` |
|
||||
| `/test distillation` | DMD distillation tests | `distillation_dmd` |
|
||||
| `/test self-forcing` | Self-Forcing tests | `self_forcing` |
|
||||
| `/test vsa` | VSA training tests | `training_vsa` |
|
||||
| `/test vmoba` | VMoBA inference tests | `inference_vmoba` |
|
||||
| `/test performance` | Performance benchmarks | `performance` |
|
||||
| `/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`) |
|
||||
|
||||
---
|
||||
|
||||
## Auto Branch Cleanup
|
||||
|
||||
After a PR is squash-merged to `main`, Mergify automatically deletes the head branch.
|
||||
Protected branches (`main`, `master`, `release/*`) are never deleted.
|
||||
|
||||
---
|
||||
|
||||
## Workflow File Reference
|
||||
|
||||
| Filename | Trigger | What it does |
|
||||
|----------|---------|-------------|
|
||||
| `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 |
|
||||
| `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 |
|
||||
| `infra-build-image.yml` | Manual (`workflow_dispatch`) | Builds Docker images for CI |
|
||||
| `infra-docs.yml` | Changes to `docs/` merged to `main` | Builds and deploys documentation to GitHub Pages |
|
||||
| `publish-fastvideo.yml` | Version bump | Publishes `fastvideo` package to PyPI |
|
||||
| `publish-kernel.yml` | Version bump | Publishes `fastvideo-kernel` package to PyPI |
|
||||
| `publish-comfyui.yml` | Version bump | Publishes ComfyUI node package to PyPI |
|
||||
@@ -1,195 +1,65 @@
|
||||
# RunPod Development Environment
|
||||
|
||||
RunPod gives you on-demand cloud GPUs for FastVideo development. It's useful when you need a beefy GPU to test training runs, benchmark inference, or reproduce results without waiting for shared cluster time.
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
## Prerequisites
|
||||
You can easily use the FastVideo Pod Template on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
- A [RunPod](https://www.runpod.io) account with billing configured
|
||||
- An SSH key pair. If you don't have one, generate it with `ssh-keygen -t ed25519`
|
||||
- Your public key (`~/.ssh/id_ed25519.pub`) ready to paste into RunPod
|
||||
|
||||
## Step 1: Create a Pod
|
||||
|
||||
**1. Verify your account**
|
||||
|
||||
Make sure you're logged into the right RunPod account before spending credits.
|
||||
## Creating a new pod
|
||||
|
||||
- Make sure you are using the correct RunPod account.
|
||||

|
||||
|
||||
**2. Filter by CUDA version**
|
||||
|
||||
Use "Additional Filters" to select CUDA 12.8.
|
||||
|
||||
- Use "Additional Filters" to select CUDA 12.8.
|
||||

|
||||
|
||||
**3. Select a GPU**
|
||||
|
||||
Click "Deploy" and pick a GPU. See [GPU Recommendations](#gpu-recommendations) below for guidance on which GPU to choose.
|
||||
|
||||
- Click "Deploy" and Pick a single A40 or RTX 4090 GPU.
|
||||

|
||||
|
||||
**4. Pick the FastVideo template**
|
||||
|
||||
Select the "FastVideo" or "fastvideo-dev" Pod Template. This pulls the pre-built image that includes all dependencies, Flash Attention, and a ready-to-use `uv` environment.
|
||||
|
||||
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
|
||||

|
||||
|
||||
**5. Name your pod**
|
||||
|
||||
Use a memorable name like `yourname-fastvideo-2026-03-28`. This helps if you have multiple pods running.
|
||||
|
||||
**6. Add a persistent volume (recommended)**
|
||||
|
||||
Attach a network volume to `/root/.cache` or `/models` for storing downloaded model weights. Models can be 10-50 GB each, and re-downloading them every session wastes time and bandwidth.
|
||||
|
||||
**7. Deploy**
|
||||
|
||||
Click Deploy. The pod takes a few minutes to start while the image pulls. You'll see it transition to "Running" in your dashboard.
|
||||
|
||||
## Step 2: Connect via SSH
|
||||
|
||||
Once the pod is running, find the "SSH over exposed TCP" connection string in the pod dashboard.
|
||||
- Set the Pod name to "`<name>-<FastVideo>-<date>`".
|
||||
|
||||
- Finally, once the pod is deployed (will take a few minutes as the image is being pulled), you can SSH into it using "SSH exposed over TCP". You'll need to use the matching private ssh key you provided.
|
||||

|
||||
|
||||
Connect with:
|
||||
## Working with the pod
|
||||
|
||||
After SSH'ing into your pod, you'll find the correct `uv` environment already activated and you should be in /FastVideo/ directory. Make sure to use /FastVideo/ for all your work.
|
||||
|
||||
To pull in the latest changes from the GitHub repo:
|
||||
|
||||
```bash
|
||||
ssh root@<pod-ip> -p <port> -i ~/.ssh/id_ed25519
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
RunPod also supports VS Code Remote SSH if you prefer an IDE.
|
||||
Run your development workflows as usual:
|
||||
|
||||
### Custom template (advanced)
|
||||
```bash
|
||||
# Run linters
|
||||
pre-commit run --all-files
|
||||
|
||||
If you're setting up a pod from scratch instead of the FastVideo template, use this image:
|
||||
# Run tests
|
||||
pytest tests/
|
||||
```
|
||||
|
||||
Make sure to push your changes back to the GitHub repo as nothing will be saved to the pod when it is terminated.
|
||||
|
||||
After you are done with your work, you can terminate the pod by clicking the "Terminate" and "Delete" buttons. Remember if the pod is not completely deleted, Runpod will keep charging you for it.
|
||||
|
||||
## Extra Information:
|
||||
If you need to customize the pod template this section has some useful information. For the most part you can leave the defaults of the FastVideo Pod Template.
|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
And paste this as the Container Start Command to enable SSH ([RunPod docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
```bash
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

|
||||
|
||||
## Step 3: Set Up FastVideo
|
||||
|
||||
After SSH'ing in, the `uv` virtual environment at `/opt/venv` is already activated (configured in `.bashrc` and `.profile`). You land in the `/FastVideo` directory.
|
||||
|
||||
**Clone or pull the repo**
|
||||
|
||||
If the pod already has the FastVideo repo:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
If starting from a blank pod:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git /FastVideo
|
||||
cd /FastVideo
|
||||
```
|
||||
|
||||
**Install the package**
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
```
|
||||
|
||||
The Docker image already includes Flash Attention and most heavy dependencies, so this is fast.
|
||||
|
||||
**Build the custom kernels (optional)**
|
||||
|
||||
VSA and STA attention kernels aren't in the Docker image by default. Build them if you're working on attention backends or need maximum inference performance:
|
||||
|
||||
```bash
|
||||
cd /FastVideo/fastvideo-kernel
|
||||
./build.sh
|
||||
```
|
||||
|
||||
The build script detects your GPU architecture automatically. An A100 or H100 takes about 5-10 minutes.
|
||||
|
||||
**Verify the setup**
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
python -c "import fastvideo; print('OK')"
|
||||
pytest tests/ -q --no-header
|
||||
```
|
||||
|
||||
## Development Workflow
|
||||
|
||||
### Editing code on RunPod
|
||||
|
||||
Two common approaches:
|
||||
|
||||
**Option A: Edit on RunPod directly**
|
||||
|
||||
Use VS Code Remote SSH or `vim`/`nano` on the pod. Commit and push when you're ready:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git add .
|
||||
git commit -m "your change"
|
||||
git push
|
||||
```
|
||||
|
||||
**Option B: Edit locally, sync to RunPod**
|
||||
|
||||
Work in your local repo, then pull on the pod:
|
||||
|
||||
```bash
|
||||
# On RunPod:
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
This keeps your local tools (editor, linters) intact while running GPU workloads on the pod.
|
||||
|
||||
### Running linters and tests
|
||||
|
||||
```bash
|
||||
# Lint
|
||||
pre-commit run --all-files
|
||||
|
||||
# Full test suite
|
||||
pytest tests/
|
||||
|
||||
# Just package tests
|
||||
pytest fastvideo/tests/ -v
|
||||
```
|
||||
|
||||
### Storing models
|
||||
|
||||
If you attached a persistent volume, point your model downloads there:
|
||||
|
||||
```bash
|
||||
export HF_HOME=/models/huggingface
|
||||
export TRANSFORMERS_CACHE=/models/huggingface
|
||||
```
|
||||
|
||||
Add these to `/root/.bashrc` so they persist across SSH sessions. The volume survives pod termination, so you only download models once.
|
||||
|
||||
### Terminating the pod
|
||||
|
||||
When you're done, push any commits you want to keep. RunPod does not save pod storage after termination.
|
||||
|
||||
Go to your RunPod dashboard, click "Terminate", then "Delete". A pod that's stopped but not deleted still charges you for storage. Fully delete it to stop all charges.
|
||||
|
||||
## GPU Recommendations
|
||||
|
||||
| GPU | VRAM | Good for |
|
||||
|-----|------|----------|
|
||||
| RTX 4090 | 24 GB | Inference testing, small model fine-tuning, quick iteration |
|
||||
| A40 | 48 GB | Mid-size training runs, 480p video generation |
|
||||
| A100 (40 GB) | 40 GB | Multi-GPU inference, training with sequence parallelism |
|
||||
| A100 (80 GB) | 80 GB | Large model training, 720p+ video generation |
|
||||
| H100 | 80 GB | Heavy training, benchmarking, kernel development |
|
||||
|
||||
For most development work, a single RTX 4090 or A40 is sufficient and cost-effective. Use A100/H100 when you need to reproduce training results at scale or test multi-GPU features.
|
||||
|
||||
@@ -1,218 +0,0 @@
|
||||
# Contributing via Pull Requests
|
||||
|
||||
This guide walks through the PR workflow: title format, labels, CI pipeline, and getting
|
||||
your changes merged.
|
||||
|
||||
---
|
||||
|
||||
## PR Title Format (Required)
|
||||
|
||||
Every PR targeting `main` must start with a type tag in square brackets. This is checked by
|
||||
Mergify before any merge is allowed.
|
||||
|
||||
**Format:**
|
||||
|
||||
```
|
||||
[type] Short description of the change
|
||||
```
|
||||
|
||||
**Valid type tags:**
|
||||
|
||||
| Tag | When to use |
|
||||
|-----|-------------|
|
||||
| `[feat]` or `[feature]` | New feature or capability |
|
||||
| `[bugfix]` or `[fix]` | Bug fix |
|
||||
| `[refactor]` | Code restructuring with no behavior change |
|
||||
| `[perf]` | Performance improvement |
|
||||
| `[ci]` | CI/CD or build tooling changes |
|
||||
| `[doc]` or `[docs]` | Documentation only |
|
||||
| `[misc]` or `[chore]` | Housekeeping, dependency bumps, minor cleanup |
|
||||
| `[kernel]` | CUDA kernel changes in `fastvideo-kernel/` |
|
||||
| `[new-model]` | Adding a new model or pipeline |
|
||||
|
||||
**Examples:**
|
||||
|
||||
```
|
||||
[feat] Add causal Wan 2.2 I2V pipeline
|
||||
[bugfix] Fix VAE temporal tiling corruption on H100
|
||||
[refactor] Restructure distributed attention dispatch
|
||||
[docs] Add LoRA finetuning guide
|
||||
[new-model] Port HunyuanVideo 1.5 to FastVideo
|
||||
```
|
||||
|
||||
If your title is missing the tag, Mergify will post a comment listing the valid formats.
|
||||
Update the title and the check will re-evaluate automatically.
|
||||
|
||||
---
|
||||
|
||||
## Labels
|
||||
|
||||
Labels are applied automatically based on your PR title and the files you changed. You don't
|
||||
need to set them manually.
|
||||
|
||||
**Type label** — set from the `[tag]` in your title:
|
||||
`type: feat`, `type: bugfix`, `type: refactor`, `type: perf`, `type: ci`, `type: docs`,
|
||||
`type: misc`, `type: new-model`
|
||||
|
||||
**Scope labels** — set from which files you modified (multiple labels can apply):
|
||||
`scope: training`, `scope: inference`, `scope: attention`, `scope: kernel`, `scope: data`,
|
||||
`scope: infra`, `scope: distributed`, `scope: docs`, `scope: ui`, `scope: model`
|
||||
|
||||
**Process labels** — set during review and merge:
|
||||
|
||||
| Label | Who sets it | Meaning |
|
||||
|-------|-------------|---------|
|
||||
| `ready` | You (`/merge` comment) or a maintainer | Triggers Full Suite and enables auto-merge |
|
||||
| `needs-rebase` | Mergify (automatic) | Your PR has conflicts; rebase against `main` |
|
||||
| `do-not-merge` | Maintainer | Blocks merge regardless of CI status |
|
||||
|
||||
---
|
||||
|
||||
## CI Pipeline
|
||||
|
||||
Three tiers run automatically on every PR.
|
||||
|
||||
**Tier 1: Pre-commit (~2 min) — runs on every push**
|
||||
|
||||
GitHub Actions checks formatting, linting, type correctness, and spelling using pre-commit
|
||||
hooks: yapf, ruff, mypy, codespell, pymarkdown, actionlint, and check-filenames.
|
||||
|
||||
**Tier 2: Fastcheck (~10-20 min) — runs on every push, path-filtered**
|
||||
|
||||
Buildkite runs GPU tests only for the components you changed. If you only modified
|
||||
`fastvideo/models/vaes/`, only VAE Tests run. Tests run in parallel.
|
||||
|
||||
**Tier 3: Full Suite (~60-90 min) — triggered by the `ready` label**
|
||||
|
||||
When you comment `/merge` (or a maintainer adds the `ready` label), Buildkite runs the
|
||||
complete test suite on your PR branch: SSIM regression, LoRA inference and training,
|
||||
distillation, self-forcing, VSA, VMoBA, performance benchmarks, and API server tests.
|
||||
|
||||
---
|
||||
|
||||
## Getting Your PR Merged
|
||||
|
||||
**Step-by-step:**
|
||||
|
||||
1. Open a PR with a title that starts with a valid `[type]` tag.
|
||||
2. Push your changes. Pre-commit and Fastcheck run automatically.
|
||||
3. Fix any pre-commit failures locally (`pre-commit run --all-files`) and push again.
|
||||
4. Wait for at least one approving review.
|
||||
5. Once approved and pre-commit is green, comment `/merge` on the PR.
|
||||
6. The `ready` label is added, which triggers the Full Suite on your PR branch.
|
||||
7. Mergify also auto-rebases your branch against `main` if it is behind and conflict-free.
|
||||
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.
|
||||
|
||||
!!! note
|
||||
Only contributors with write permission to the repository can trigger slash commands.
|
||||
If you're an external contributor, ask a maintainer to run `/merge` or add the `ready`
|
||||
label for you.
|
||||
|
||||
---
|
||||
|
||||
## Running Tests On Demand
|
||||
|
||||
Comment on your PR to trigger specific tests independently of the auto-merge flow.
|
||||
|
||||
**Trigger the entire Full Suite:**
|
||||
|
||||
```
|
||||
/test full
|
||||
```
|
||||
|
||||
**Trigger the Fastcheck suite:**
|
||||
|
||||
```
|
||||
/test fastcheck
|
||||
```
|
||||
|
||||
**Trigger individual tests:**
|
||||
|
||||
```
|
||||
/test encoder # Encoder component tests
|
||||
/test vae # VAE component tests
|
||||
/test transformer # Transformer / DiT tests
|
||||
/test kernel # CUDA kernel tests
|
||||
/test unit # Unit tests
|
||||
|
||||
/test ssim # SSIM regression tests
|
||||
/test training # Training pipeline tests
|
||||
/test lora-inference # LoRA inference tests
|
||||
/test lora-training # LoRA training tests
|
||||
/test distillation # DMD distillation tests
|
||||
/test self-forcing # Self-Forcing distillation tests
|
||||
/test vsa # VSA training tests
|
||||
/test vmoba # VMoBA inference tests
|
||||
/test performance # Performance benchmarks
|
||||
/test api # API server integration tests
|
||||
```
|
||||
|
||||
The workflow reacts with a 🚀 emoji to confirm the command was received.
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Pre-commit fails
|
||||
|
||||
Run locally to reproduce and auto-fix:
|
||||
|
||||
```bash
|
||||
# Install pre-commit if needed
|
||||
uv pip install pre-commit
|
||||
pre-commit install
|
||||
|
||||
# Run all checks on all files
|
||||
pre-commit run --all-files
|
||||
```
|
||||
|
||||
Common quick fixes:
|
||||
|
||||
- **yapf**: `yapf -i <file>` (Python formatting)
|
||||
- **ruff**: `ruff check --fix <file>` (linting)
|
||||
- **codespell**: `codespell --write-changes <file>` (spelling)
|
||||
|
||||
### PR title format check fails
|
||||
|
||||
Update your title to start with a valid type tag. The Mergify merge protection check
|
||||
re-evaluates automatically after you save the title.
|
||||
|
||||
Valid tags: `feat`, `feature`, `bugfix`, `fix`, `refactor`, `perf`, `ci`, `doc`, `docs`,
|
||||
`misc`, `chore`, `kernel`, `new-model`
|
||||
|
||||
### My PR has merge conflicts (`needs-rebase` label)
|
||||
|
||||
Rebase against `main` and force-push:
|
||||
|
||||
```bash
|
||||
git fetch origin main
|
||||
git rebase origin/main
|
||||
# Resolve any conflicts, then:
|
||||
git push --force-with-lease
|
||||
```
|
||||
|
||||
Mergify removes the `needs-rebase` label automatically once conflicts are resolved.
|
||||
|
||||
### 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.
|
||||
|
||||
Common causes:
|
||||
|
||||
- Test failures caused by your code changes
|
||||
- Missing dependency in `pyproject.toml`
|
||||
- GPU memory issue (some tests require specific hardware like L40S or H100)
|
||||
- Kernel build failure (if you changed `fastvideo-kernel/`)
|
||||
|
||||
After fixing, push and comment `/merge` again.
|
||||
|
||||
### I'm an external contributor without write permission
|
||||
|
||||
You can't use slash commands directly. After your PR is approved, ask a maintainer to
|
||||
comment `/merge` or add the `ready` label.
|
||||
+22
-109
@@ -23,18 +23,13 @@ SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos us
|
||||
|
||||
```
|
||||
fastvideo/tests/ssim/
|
||||
├── reference_videos/
|
||||
│ ├── default/
|
||||
│ │ └── <GPU>_reference_videos/
|
||||
│ │ ├── <Model_Name>/
|
||||
│ │ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
|
||||
│ │ │ │ └── <Video_File>
|
||||
│ └── full_quality/
|
||||
│ └── <GPU>_reference_videos/
|
||||
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
|
||||
│ ├── <Model_Name>/
|
||||
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
|
||||
│ │ │ └── <Video_File>
|
||||
├── test_causal_similarity.py
|
||||
├── test_wan_t2v_similarity.py
|
||||
├── test_wan_i2v_similarity.py
|
||||
├── reference_videos_cli.py
|
||||
├── test_inference_similarity.py
|
||||
├── update_reference_videos.sh
|
||||
└── ...
|
||||
```
|
||||
|
||||
@@ -42,7 +37,7 @@ fastvideo/tests/ssim/
|
||||
|
||||
To add a new SSIM test, follow these steps:
|
||||
|
||||
1. **Create or Update a Test File**: Prefer model-specific files (for example `test_wan_t2v_similarity.py`) and create a new one when testing a distinct model or pipeline.
|
||||
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
|
||||
|
||||
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
|
||||
|
||||
@@ -86,15 +81,10 @@ To add a new SSIM test, follow these steps:
|
||||
```
|
||||
|
||||
4. **Reference Videos**:
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos/<quality-tier>/<GPU>_reference_videos`.
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
|
||||
* Inspect the generated video to ensure it meets quality expectations.
|
||||
* Move the generated video to the appropriate quality/GPU reference folder:
|
||||
`fastvideo/tests/ssim/reference_videos/<quality-tier>/<GPU>_reference_videos/<Model>/<Backend>/`.
|
||||
* You can use the helper CLI to copy generated videos into a reference folder:
|
||||
`python fastvideo/tests/ssim/reference_videos_cli.py copy-local --quality-tier default --reference-dir fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos`
|
||||
* Upload/download can target both quality tiers and specific GPU folders:
|
||||
`python fastvideo/tests/ssim/reference_videos_cli.py upload --quality-tier all`
|
||||
`python fastvideo/tests/ssim/reference_videos_cli.py download --quality-tier full_quality --device-folder H200_reference_videos`
|
||||
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
|
||||
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
|
||||
|
||||
### Running Tests Locally
|
||||
|
||||
@@ -106,113 +96,36 @@ pytest fastvideo/tests/ssim/ -vs
|
||||
|
||||
Ensure you have the necessary GPUs available as defined in your test parameters.
|
||||
|
||||
## CI Integration
|
||||
## Modal Workflow
|
||||
|
||||
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The
|
||||
workflow scripts are located in `fastvideo/tests/modal/`.
|
||||
|
||||
### Buildkite Pipeline
|
||||
|
||||
Tests are orchestrated by Buildkite (`.buildkite/pipeline.yml`) and executed on Modal GPU
|
||||
instances. The pipeline runs in two modes:
|
||||
|
||||
**Fastcheck** — runs on every PR push, path-filtered. Only tests for the components you
|
||||
changed are triggered. Tests run in parallel.
|
||||
|
||||
**Full Suite** — runs when a PR enters the Merge Queue (or when triggered manually via
|
||||
`/test full`). Covers SSIM regression, training, distillation, inference, and performance.
|
||||
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
|
||||
|
||||
### `pr_test.py`
|
||||
|
||||
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines
|
||||
Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
|
||||
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
|
||||
|
||||
### Updating Modal Configuration
|
||||
|
||||
If you add a new test that requires:
|
||||
|
||||
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
|
||||
* **Longer Execution Time**: Increase the `timeout` parameter.
|
||||
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image
|
||||
environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY`
|
||||
is passed.
|
||||
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
|
||||
|
||||
For SSIM tests, use `fastvideo/tests/modal/ssim_test.py`:
|
||||
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
|
||||
|
||||
```bash
|
||||
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests
|
||||
```python
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
```
|
||||
|
||||
Target specific SSIM files/models:
|
||||
|
||||
```bash
|
||||
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests \
|
||||
--test-files test_wan_t2v_similarity.py \
|
||||
--model-ids Wan2.1-T2V-1.3B-Diffusers
|
||||
```
|
||||
|
||||
If HF token env vars are not set (`HF_API_KEY` / `HUGGINGFACE_HUB_TOKEN` /
|
||||
`HF_TOKEN`), the local entrypoint fails fast. To export raw `generated_videos`
|
||||
from Modal to the shared volume:
|
||||
|
||||
```bash
|
||||
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests \
|
||||
--sync-generated-to-volume
|
||||
```
|
||||
|
||||
The raw export path is quality-tiered:
|
||||
|
||||
* default params: `ssim_generated_videos/default/<subdir>/generated_videos`
|
||||
* full-quality params: `ssim_generated_videos/full_quality/<subdir>/generated_videos`
|
||||
|
||||
The printed `modal volume get` command also downloads into a quality-specific
|
||||
local directory under `./generated_videos_modal/<quality-tier>`.
|
||||
|
||||
To turn downloaded Modal outputs into local reference videos, use the matching
|
||||
quality tier with `copy-local`, for example:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--quality-tier full_quality \
|
||||
--generated-dir ./generated_videos_modal/full_quality/L40S_reference_videos \
|
||||
--device-folder L40S_reference_videos
|
||||
```
|
||||
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
|
||||
|
||||
### Workflow Scripts
|
||||
|
||||
The shell script that triggers tests in CI is `.buildkite/scripts/pr_test.sh`. If you add
|
||||
a new test category (e.g., a new folder outside of `ssim`), you will need to:
|
||||
|
||||
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
|
||||
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
|
||||
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
|
||||
|
||||
!!! note
|
||||
If you are a maintainer, update the workflow script in Buildkite after merging. Otherwise,
|
||||
ask a maintainer for help.
|
||||
|
||||
## Triggering Tests via Slash Commands
|
||||
|
||||
Maintainers and contributors with write access can trigger individual test suites directly
|
||||
from a PR comment. The workflow reacts with a 🚀 emoji to confirm the command was received.
|
||||
|
||||
```
|
||||
/test ssim # SSIM regression tests
|
||||
/test training # Training pipeline tests
|
||||
/test lora-training # LoRA training tests
|
||||
/test lora-inference # LoRA inference tests
|
||||
/test distillation # DMD distillation tests
|
||||
/test self-forcing # Self-Forcing tests
|
||||
/test vsa # VSA training tests
|
||||
/test vmoba # VMoBA inference tests
|
||||
/test performance # Performance benchmarks
|
||||
/test api # API server integration tests
|
||||
/test encoder # Encoder component tests (Fastcheck)
|
||||
/test vae # VAE component tests (Fastcheck)
|
||||
/test transformer # Transformer / DiT tests (Fastcheck)
|
||||
/test kernel # CUDA kernel tests (Fastcheck)
|
||||
/test unit # Unit tests (Fastcheck)
|
||||
/test full # Entire Full Suite
|
||||
/test fastcheck # Entire Fastcheck suite
|
||||
```
|
||||
|
||||
See [CI Architecture](ci_architecture.md) for the complete reference.
|
||||
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
|
||||
|
||||
+60
-34
@@ -42,7 +42,8 @@ def fix_case(text: str) -> str:
|
||||
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
|
||||
}
|
||||
for pattern, repl in subs.items():
|
||||
text = re.sub(rf'\b{pattern}\b', repl, text, flags=re.IGNORECASE) # type: ignore[call-overload]
|
||||
text = re.sub(rf'\b{pattern}\b', repl, text,
|
||||
flags=re.IGNORECASE) # type: ignore[call-overload]
|
||||
return text
|
||||
|
||||
|
||||
@@ -133,7 +134,9 @@ class Example:
|
||||
if not markdown_files:
|
||||
raise IndexError(f"No Markdown files found in {self.path}")
|
||||
|
||||
readme_files = [f for f in markdown_files if f.name.lower() == "readme.md"]
|
||||
readme_files = [
|
||||
f for f in markdown_files if f.name.lower() == "readme.md"
|
||||
]
|
||||
if readme_files:
|
||||
return readme_files[0]
|
||||
|
||||
@@ -153,7 +156,8 @@ class Example:
|
||||
if self.path.is_file():
|
||||
return []
|
||||
is_other_file = lambda file: file.is_file() and file != self.main_file
|
||||
return [file for file in self.path.rglob("*") if is_other_file(file)] # type: ignore[no-untyped-call]
|
||||
return [file for file in self.path.rglob("*")
|
||||
if is_other_file(file)] # type: ignore[no-untyped-call]
|
||||
|
||||
def determine_title(self) -> str:
|
||||
return fix_case(self.path.stem.replace("_", " ").title())
|
||||
@@ -186,8 +190,8 @@ class Example:
|
||||
content += "## Additional Files\n\n"
|
||||
# Define binary/non-text file extensions to skip
|
||||
binary_extensions = {
|
||||
'.mp4', '.avi', '.mov', '.mkv', '.gif', '.jpg', '.jpeg', '.png', '.webp', '.bmp', '.pdf', '.zip', '.tar',
|
||||
'.gz', '.mp3', '.wav'
|
||||
'.mp4', '.avi', '.mov', '.mkv', '.gif', '.jpg', '.jpeg', '.png',
|
||||
'.webp', '.bmp', '.pdf', '.zip', '.tar', '.gz', '.mp3', '.wav'
|
||||
}
|
||||
|
||||
for file in sorted(self.other_files):
|
||||
@@ -255,7 +259,8 @@ def create_category_indices() -> dict[str, Index]:
|
||||
category_indices = {
|
||||
"inference":
|
||||
Index(
|
||||
path=ROOT_DIR / "docs/inference/examples/examples_inference_index.md",
|
||||
path=ROOT_DIR /
|
||||
"docs/inference/examples/examples_inference_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with [basic.md](basic.md).",
|
||||
@@ -266,15 +271,18 @@ def create_category_indices() -> dict[str, Index]:
|
||||
Index(
|
||||
path=ROOT_DIR / "docs/training/examples/examples_training_index.md",
|
||||
title="🚀 Examples",
|
||||
description="Training examples demonstrate how to use FastVideo training.",
|
||||
description=
|
||||
"Training examples demonstrate how to use FastVideo training.",
|
||||
caption="Examples",
|
||||
maxdepth=3,
|
||||
),
|
||||
"distillation":
|
||||
Index(
|
||||
path=ROOT_DIR / "docs/distillation/examples/examples_distillation_index.md",
|
||||
path=ROOT_DIR /
|
||||
"docs/distillation/examples/examples_distillation_index.md",
|
||||
title="🚀 Examples",
|
||||
description="Distillation examples demonstrate how to use FastVideo distillation.",
|
||||
description=
|
||||
"Distillation examples demonstrate how to use FastVideo distillation.",
|
||||
caption="Examples",
|
||||
maxdepth=3,
|
||||
),
|
||||
@@ -288,7 +296,8 @@ def create_category_indices() -> dict[str, Index]:
|
||||
return category_indices
|
||||
|
||||
|
||||
def find_examples(category_indices: dict[str, Index], generate_main_index: bool) -> list[Example]:
|
||||
def find_examples(category_indices: dict[str, Index],
|
||||
generate_main_index: bool) -> list[Example]:
|
||||
"""Find all examples from the examples directory."""
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
@@ -330,9 +339,13 @@ def find_examples(category_indices: dict[str, Index], generate_main_index: bool)
|
||||
return examples
|
||||
|
||||
|
||||
def create_nested_structures(examples: list[Example]) -> dict[str, dict[str, dict[str, dict[str, NestedStructure]]]]:
|
||||
def create_nested_structures(
|
||||
examples: list[Example]
|
||||
) -> dict[str, dict[str, dict[str, dict[str, NestedStructure]]]]:
|
||||
"""Create nested structures for training and distillation categories."""
|
||||
nested_structures: dict[str, dict[str, dict[str, dict[str, NestedStructure]]]] = {}
|
||||
nested_structures: dict[str, dict[str, dict[str,
|
||||
dict[str,
|
||||
NestedStructure]]]] = {}
|
||||
|
||||
# Map category names to actual directory names
|
||||
category_dir_mapping = {
|
||||
@@ -365,11 +378,13 @@ def create_nested_structures(examples: list[Example]) -> dict[str, dict[str, dic
|
||||
nested_structures[example.category][method][model] = {}
|
||||
|
||||
# Store the nested structure
|
||||
nested_structures[example.category][method][model][dataset] = NestedStructure(category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
nested_structures[
|
||||
example.category][method][model][dataset] = NestedStructure(
|
||||
category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
|
||||
elif example.category == "distillation" and len(path_parts) >= 2:
|
||||
# For distillation examples like Wan2.1-T2V/Wan-Syn-Data-480P
|
||||
@@ -386,16 +401,20 @@ def create_nested_structures(examples: list[Example]) -> dict[str, dict[str, dic
|
||||
nested_structures[example.category][method][model] = {}
|
||||
|
||||
# Store the nested structure
|
||||
nested_structures[example.category][method][model][dataset] = NestedStructure(category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
nested_structures[
|
||||
example.category][method][model][dataset] = NestedStructure(
|
||||
category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
|
||||
return nested_structures
|
||||
|
||||
|
||||
def generate_flat_examples(examples: list[Example], category_indices: dict[str, Index], examples_index: Index | None,
|
||||
def generate_flat_examples(examples: list[Example],
|
||||
category_indices: dict[str, Index],
|
||||
examples_index: Index | None,
|
||||
generate_main_index: bool) -> None:
|
||||
"""Generate documentation for flat structure examples (inference, etc.)."""
|
||||
for example in examples:
|
||||
@@ -418,8 +437,9 @@ def generate_flat_examples(examples: list[Example], category_indices: dict[str,
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
|
||||
def generate_nested_examples(nested_structures: dict[str, dict[str, dict[str, dict[str, NestedStructure]]]],
|
||||
category_indices: dict[str, Index]) -> None:
|
||||
def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
|
||||
str, dict[str, NestedStructure]]]], category_indices: dict[str,
|
||||
Index]) -> None:
|
||||
"""Generate documentation for nested structure examples (training, distillation)."""
|
||||
for category_name in ["training", "distillation"]:
|
||||
if category_name not in category_indices or category_name not in nested_structures:
|
||||
@@ -444,11 +464,12 @@ def generate_nested_examples(nested_structures: dict[str, dict[str, dict[str, di
|
||||
f.write(nested_struct.example.generate())
|
||||
|
||||
# Create model-level index
|
||||
model_index = Index(path=category_base_dir / f"{model}.md",
|
||||
title=fix_case(model.replace('_', ' ')),
|
||||
description=f"Examples for the {model} model.",
|
||||
caption=f"{fix_case(model.replace('_', ' '))} Datasets",
|
||||
maxdepth=1)
|
||||
model_index = Index(
|
||||
path=category_base_dir / f"{model}.md",
|
||||
title=fix_case(model.replace('_', ' ')),
|
||||
description=f"Examples for the {model} model.",
|
||||
caption=f"{fix_case(model.replace('_', ' '))} Datasets",
|
||||
maxdepth=1)
|
||||
|
||||
# Add dataset indices to model index
|
||||
for dataset, nested_struct in datasets.items():
|
||||
@@ -487,7 +508,8 @@ def generate_examples(generate_main_index: bool = False) -> None:
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description="A collection of examples demonstrating usage of FastVideo.\n\n"
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\n\n"
|
||||
f"All documented examples are autogenerated using [generate_examples.py](https://github.com/{GITHUB_REPO}/blob/main/docs/generate_examples.py) "
|
||||
f"from examples found in the [examples](https://github.com/{GITHUB_REPO}/tree/main/examples) directory.",
|
||||
caption="Examples",
|
||||
@@ -500,7 +522,8 @@ def generate_examples(generate_main_index: bool = False) -> None:
|
||||
nested_structures = create_nested_structures(examples)
|
||||
|
||||
# Generate flat structure examples (inference, etc.)
|
||||
generate_flat_examples(examples, category_indices, examples_index, generate_main_index)
|
||||
generate_flat_examples(examples, category_indices, examples_index,
|
||||
generate_main_index)
|
||||
|
||||
# Generate nested structure examples (training, distillation)
|
||||
generate_nested_examples(nested_structures, category_indices)
|
||||
@@ -511,8 +534,11 @@ def generate_examples(generate_main_index: bool = False) -> None:
|
||||
# Add to main index if it exists
|
||||
if generate_main_index and examples_index:
|
||||
main_index_dir = examples_index.path.parent
|
||||
rel_path = os.path.relpath(category_index.path, start=main_index_dir)
|
||||
examples_index.documents.insert(0, str(rel_path).replace("\\", "/").replace(".md", ""))
|
||||
rel_path = os.path.relpath(category_index.path,
|
||||
start=main_index_dir)
|
||||
examples_index.documents.insert(
|
||||
0,
|
||||
str(rel_path).replace("\\", "/").replace(".md", ""))
|
||||
|
||||
# Write the category index file
|
||||
with open(category_index.path, "w+") as f:
|
||||
|
||||
@@ -1,93 +0,0 @@
|
||||
# Training
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Single-node
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh <config.yaml> [--dotted.key value ...]
|
||||
```
|
||||
|
||||
The script auto-detects available GPUs. Override with `NUM_GPUS`:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 bash examples/train/run.sh \
|
||||
examples/train/configs/fine_tuning/wan/t2v.yaml
|
||||
```
|
||||
|
||||
### Multi-node (Slurm)
|
||||
|
||||
```bash
|
||||
bash examples/train/run_slurm.sh <config.yaml> <num_nodes> [--dotted.key value ...]
|
||||
```
|
||||
|
||||
```bash
|
||||
bash examples/train/run_slurm.sh \
|
||||
examples/train/configs/fine_tuning/wan/t2v.yaml 4 \
|
||||
--training.distributed.hsdp_shard_dim 8 \
|
||||
--training.distributed.hsdp_replicate_dim 4
|
||||
```
|
||||
|
||||
Slurm environment variables:
|
||||
|
||||
| Variable | Default | Description |
|
||||
|---|---|---|
|
||||
| `PARTITION` | `main` | Slurm partition |
|
||||
| `NUM_GPUS` | `8` | GPUs per node |
|
||||
| `CPUS_PER_TASK` | `128` | CPUs per task |
|
||||
| `MEM` | `1440G` | Memory per node |
|
||||
| `EXCLUDE` | `""` | Nodes to exclude |
|
||||
|
||||
### Overriding config values
|
||||
|
||||
Any config field can be overridden from the command line using dotted keys:
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh examples/train/configs/fine_tuning/wan/t2v.yaml \
|
||||
--training.optimizer.learning_rate 1e-5 \
|
||||
--training.loop.max_train_steps 1000 \
|
||||
--training.checkpoint.output_dir outputs/my_experiment
|
||||
```
|
||||
|
||||
### Resuming from a checkpoint
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh examples/train/configs/fine_tuning/wan/t2v.yaml \
|
||||
--training.checkpoint.resume_from_checkpoint outputs/my_experiment/checkpoint-500
|
||||
```
|
||||
|
||||
## W&B Logging
|
||||
|
||||
Training metrics and validation videos are logged to
|
||||
[Weights & Biases](https://wandb.ai). Set your API key before launching:
|
||||
|
||||
```bash
|
||||
export WANDB_API_KEY=your_key_here
|
||||
```
|
||||
|
||||
To disable W&B (e.g. for local debugging):
|
||||
|
||||
```bash
|
||||
export WANDB_MODE=offline
|
||||
```
|
||||
|
||||
The project name and run name are set in the config under `training.tracker`:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
## Directory Layout
|
||||
|
||||
```
|
||||
examples/train/
|
||||
├── configs/ # Single-step training configs (by method and model)
|
||||
├── scenario/ # Multi-step end-to-end training pipelines
|
||||
├── run.sh # Single-node launcher
|
||||
└── run_slurm.sh # Multi-node Slurm launcher
|
||||
```
|
||||
|
||||
See `configs/README.md` and `scenario/README.md` for details.
|
||||
@@ -1,21 +0,0 @@
|
||||
# Training Configs
|
||||
|
||||
Single-step training configurations organized by method and model.
|
||||
|
||||
```
|
||||
configs/
|
||||
├── fine_tuning/ # Standard finetuning and DFSFT
|
||||
├── distribution_matching/ # DMD2 and Self-Forcing
|
||||
├── knowledge_distillation/ # KD from teacher to student
|
||||
└── example.yaml # Annotated reference config with all fields
|
||||
```
|
||||
|
||||
Each method directory contains per-model subdirectories (e.g. `wan/`, `hunyuan/`).
|
||||
|
||||
Launch any config with:
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh examples/train/configs/<method>/<model>/<config>.yaml
|
||||
```
|
||||
|
||||
For multi-step training pipelines, see `examples/train/scenario/`.
|
||||
@@ -0,0 +1,95 @@
|
||||
# V3 config: WanGame causal Diffusion-Forcing SFT (DFSFT).
|
||||
#
|
||||
# Uses _target_-based instantiation — each model role is an independent
|
||||
# class instance; the method class is resolved directly from the YAML.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wangame.WanGameCausalModel
|
||||
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
|
||||
trainable: true
|
||||
# transformer_override_safetensor: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500/transformer/model.safetensors
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
attn_kind: dense
|
||||
# use_ema: true
|
||||
chunk_size: 3
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: >-
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 352
|
||||
num_width: 640
|
||||
num_frames: 69
|
||||
apply_bot_died_filter: true
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1e-4
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 1e-5
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wangame_dfsft_causal_v3
|
||||
# resume_from_checkpoint: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500
|
||||
training_state_checkpointing_steps: 1000000
|
||||
weight_only_checkpointing_steps: 1000000
|
||||
checkpoints_total_limit: 0
|
||||
best_checkpoint_start_step: 1000000
|
||||
best_checkpoint_top_k: 0
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wangame_r
|
||||
run_name: wangame_dfsft_causal_v3
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
# ema:
|
||||
# _target_: fastvideo.train.callbacks.ema.EMACallback
|
||||
# beta: 0.9999
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalOdeDMDPipeline
|
||||
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_8.json
|
||||
every_steps: 500
|
||||
sampling_steps: [40]
|
||||
scheduler_target: fastvideo.models.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler
|
||||
guidance_scale: 1.0
|
||||
num_frames: 69
|
||||
evaluate_ptlflow: false
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -1,72 +0,0 @@
|
||||
# HunyuanVideo T2V finetune config.
|
||||
#
|
||||
# Data must be preprocessed with Hunyuan VAE + dual text encoders
|
||||
# (LLaMA + CLIP) into parquet format before training.
|
||||
# See scripts/preprocess/ for preprocessing pipeline.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.hunyuan.HunyuanModel
|
||||
init_from: hunyuanvideo-community/HunyuanVideo
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/hunyuan_training_data
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 32
|
||||
num_height: 720
|
||||
num_width: 1280
|
||||
num_frames: 125
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 5000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/hunyuan_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_hunyuan
|
||||
run_name: hunyuan_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.hunyuan.hunyuan_pipeline.HunyuanVideoPipeline
|
||||
dataset_file: data/hunyuan_training_data/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 7
|
||||
@@ -1,68 +0,0 @@
|
||||
# Wan 2.1 T2V 1.3B finetune.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 20
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan_finetune
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo
|
||||
run_name: wan_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline
|
||||
dataset_file: data/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 5.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -1,89 +0,0 @@
|
||||
# Knowledge Distillation: Wan 2.1 Causal T2V 1.3B
|
||||
#
|
||||
# Trains a causal student (1.3B) with per-frame block-quantized MSE loss
|
||||
# on teacher (14B) ODE trajectories. The resulting checkpoint serves as
|
||||
# the ode_init weight initialization for downstream Self-Forcing training.
|
||||
#
|
||||
# The same teacher_path_cache can be shared with kd_wan2.1_t2v_1.3B.yaml
|
||||
# since the cache stores teacher trajectories independent of the student.
|
||||
#
|
||||
# After training, export the ode_init checkpoint:
|
||||
# python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
# --role student \
|
||||
# --output_dir outputs/wan2.1_causal_kd_ode_init/diffusers \
|
||||
# --checkpoint_dir outputs/wan2.1_causal_kd_ode_init/latest
|
||||
#
|
||||
# The exported model.safetensors can then be passed as:
|
||||
# transformer_override_safetensor in a Self-Forcing config.
|
||||
#
|
||||
# Remove the 'teacher' block once the cache is complete to avoid
|
||||
# loading the 14B model on subsequent runs.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher: # remove once cache is complete
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.knowledge_distillation.kd.KDCausalMethod
|
||||
teacher_path_cache: data/kd_cache/wan14b_4step
|
||||
t_list: [999, 937, 833, 624, 0] # student inference schedule (num_steps+1)
|
||||
student_sample_steps: 4
|
||||
teacher_inference_steps: 48 # teacher ODE steps; nearest indices selected for t_list
|
||||
teacher_guidance_scale: 3.5 # CFG scale during teacher path generation
|
||||
num_frames_per_block: 3 # frames sharing the same noise level (per-frame KD)
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 2
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 20
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 7e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 10000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_kd_ode_init
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: kd-wan
|
||||
run_name: wan2.1_causal_kd_ode_init
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,85 +0,0 @@
|
||||
# Knowledge Distillation: Wan 2.1 T2V 1.3B
|
||||
#
|
||||
# Trains student (1.3B) with MSE loss on teacher (14B) ODE trajectories.
|
||||
# The resulting checkpoint serves as the ode_init weight initialization
|
||||
# for downstream Self-Forcing training.
|
||||
#
|
||||
# After training, export the ode_init checkpoint:
|
||||
# python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
# --role student \
|
||||
# --output_dir outputs/wan2.1_kd_ode_init/diffusers \
|
||||
# --checkpoint_dir outputs/wan2.1_kd_ode_init/latest
|
||||
#
|
||||
# The exported model.safetensors can then be passed as:
|
||||
# transformer_override_safetensor in a Self-Forcing config.
|
||||
#
|
||||
# Remove the 'teacher' block once the cache is complete to avoid
|
||||
# loading the 14B model on subsequent runs.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher: # remove once cache is complete
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.knowledge_distillation.kd.KDMethod
|
||||
teacher_path_cache: data/kd_cache/wan14b_4step
|
||||
t_list: [999, 937, 833, 624, 0] # student inference schedule (num_steps+1)
|
||||
student_sample_steps: 4
|
||||
teacher_inference_steps: 48 # teacher ODE steps; nearest indices selected for t_list
|
||||
teacher_guidance_scale: 3.5 # CFG scale during teacher path generation
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 2
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 20
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 7e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 10000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_kd_ode_init
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: kd-wan
|
||||
run_name: wan2.1_kd_ode_init
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,78 +0,0 @@
|
||||
# HunyuanVideo T2V overfitting test config.
|
||||
#
|
||||
# Overfits on 3 short videos (480x832, 77 frames) to verify the
|
||||
# training pipeline works end-to-end.
|
||||
#
|
||||
# Preprocess data first:
|
||||
# CUDA_VISIBLE_DEVICES=0 python scripts/preprocess_hunyuan_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# bash examples/train/run.sh examples/train/configs/overfit_hunyuan_t2v.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.hunyuan.HunyuanModel
|
||||
init_from: hunyuanvideo-community/HunyuanVideo
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/hunyuan_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 20
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/hunyuan_overfit
|
||||
training_state_checkpointing_steps: 50
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_hunyuan
|
||||
run_name: hunyuan_overfit
|
||||
|
||||
model:
|
||||
precondition_outputs: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.hunyuan.hunyuan_pipeline.HunyuanVideoPipeline
|
||||
dataset_file: data/hunyuan_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 50
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 7
|
||||
@@ -0,0 +1,102 @@
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wangame.WanGameCausalModel
|
||||
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/SFWanGame-2.1-0308-10000steps
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wangame.WanGameModel
|
||||
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wangame.WanGameModel
|
||||
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 1.0
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: >-
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 352
|
||||
num_width: 640
|
||||
num_frames: 69
|
||||
apply_bot_died_filter: true
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1e-5
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wangame_1.3b_self_forcing
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 5
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: wangame_1.3b_self_forcing
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalSdeDMDPipeline
|
||||
scheduler_target: fastvideo.models.schedulers.scheduling_self_forcing_flow_match.SelfForcingFlowMatchScheduler
|
||||
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_4.json
|
||||
every_steps: 5
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 69
|
||||
guidance_scale: 1.0
|
||||
evaluate_ptlflow: false
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,94 @@
|
||||
# V3 config: WanGame causal Diffusion-Forcing SFT (DFSFT).
|
||||
#
|
||||
# Uses _target_-based instantiation — each model role is an independent
|
||||
# class instance; the method class is resolved directly from the YAML.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wangame.WanGameCausalModel
|
||||
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
attn_kind: dense
|
||||
# use_ema: true
|
||||
chunk_size: 3
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: >-
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
|
||||
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 352
|
||||
num_width: 640
|
||||
num_frames: 69
|
||||
apply_bot_died_filter: true
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1e-5
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 10000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wangame_tfsft_causal
|
||||
# resume_from_checkpoint: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500
|
||||
training_state_checkpointing_steps: 5000
|
||||
weight_only_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 10
|
||||
best_checkpoint_start_step: 2000
|
||||
best_checkpoint_top_k: 5
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wangame_r
|
||||
run_name: wangame_tfsft_causal
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
# ema:
|
||||
# _target_: fastvideo.train.callbacks.ema.EMACallback
|
||||
# beta: 0.9999
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalOdeDMDPipeline
|
||||
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_4.json
|
||||
every_steps: 500
|
||||
sampling_steps: [40]
|
||||
scheduler_target: fastvideo.models.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler
|
||||
guidance_scale: 1.0
|
||||
num_frames: 69
|
||||
evaluate_ptlflow: false
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
+13
-6
@@ -5,12 +5,12 @@
|
||||
# bash examples/train/run.sh <config.yaml> [--dotted.key value ...]
|
||||
#
|
||||
# Examples:
|
||||
# bash examples/train/run.sh examples/train/configs/fine_tuning/wan/t2v.yaml
|
||||
# bash examples/train/run.sh examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml --dry-run
|
||||
# bash examples/train/run.sh examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml \
|
||||
# bash examples/train/run.sh examples/train/finetune_wan2.1_t2v_1.3B_vsa_phase3.4_0.9sparsity.yaml
|
||||
# bash examples/train/run.sh examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml --dry-run
|
||||
# bash examples/train/run.sh examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml \
|
||||
# --training.distributed.num_gpus 4 \
|
||||
# --training.optimizer.learning_rate 1e-5
|
||||
# bash examples/train/run.sh examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml \
|
||||
# bash examples/train/run.sh examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml \
|
||||
# --training.checkpoint.resume_from_checkpoint outputs/my_run/checkpoint-1000
|
||||
#
|
||||
# Logs are written to logs/<config_name>_<timestamp>.log (and also printed to stdout).
|
||||
@@ -22,16 +22,17 @@ shift
|
||||
|
||||
# ── GPU / node settings ──────────────────────────────────────────
|
||||
NUM_GPUS="${NUM_GPUS:-$(nvidia-smi -L 2>/dev/null | wc -l)}"
|
||||
NUM_GPUS="${NUM_GPUS:-1}"
|
||||
NUM_GPUS="${NUM_GPUS:-8}"
|
||||
NNODES="${NNODES:-1}"
|
||||
NODE_RANK="${NODE_RANK:-0}"
|
||||
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
|
||||
MASTER_PORT="${MASTER_PORT:-29501}"
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# ── W&B ──────────────────────────────────────────────────────────
|
||||
export WANDB_API_KEY="${WANDB_API_KEY:-}"
|
||||
export WANDB_API_KEY="${WANDB_API_KEY:-7ff8b6e8356924f7a6dd51a0342dd1a422ea9352}"
|
||||
export WANDB_MODE="${WANDB_MODE:-online}"
|
||||
|
||||
|
||||
# ── Log file ─────────────────────────────────────────────────────
|
||||
CONFIG_NAME="$(basename "${CONFIG}" .yaml)"
|
||||
TIMESTAMP="$(date +%Y%m%d_%H%M%S)"
|
||||
@@ -39,6 +40,12 @@ LOG_DIR="${LOG_DIR:-examples/train}"
|
||||
mkdir -p "${LOG_DIR}"
|
||||
LOG_FILE="${LOG_DIR}/${CONFIG_NAME}_${TIMESTAMP}.log"
|
||||
|
||||
set +u
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate mhuo-fv
|
||||
set -u
|
||||
export PYTHONPATH="/mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor:${PYTHONPATH:-}"
|
||||
|
||||
echo "=== Train Training ==="
|
||||
echo "Config: ${CONFIG}"
|
||||
echo "Num GPUs: ${NUM_GPUS}"
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=wg-sf
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=examples/train/slurm_%j.out
|
||||
#SBATCH --error=examples/train/slurm_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
CONFIG="${1:?Usage: sbatch examples/train/run.slurm <config.yaml> [extra flags...]}"
|
||||
shift
|
||||
EXTRA_ARGS=("$@")
|
||||
set --
|
||||
|
||||
cd /mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor
|
||||
|
||||
get_num_gpus() {
|
||||
if [[ -n "${NUM_GPUS:-}" ]]; then
|
||||
echo "${NUM_GPUS}"
|
||||
return
|
||||
fi
|
||||
if [[ -n "${SLURM_GPUS_ON_NODE:-}" ]]; then
|
||||
echo "${SLURM_GPUS_ON_NODE%%(*}"
|
||||
return
|
||||
fi
|
||||
if [[ -n "${SLURM_GPUS_PER_NODE:-}" ]]; then
|
||||
echo "${SLURM_GPUS_PER_NODE%%(*}"
|
||||
return
|
||||
fi
|
||||
if command -v nvidia-smi >/dev/null 2>&1; then
|
||||
nvidia-smi -L 2>/dev/null | wc -l | tr -d " "
|
||||
else
|
||||
echo 8
|
||||
fi
|
||||
}
|
||||
|
||||
export NNODES="${NNODES:-${SLURM_JOB_NUM_NODES:-1}}"
|
||||
export NUM_GPUS="$(get_num_gpus)"
|
||||
export MASTER_PORT="${MASTER_PORT:-29501}"
|
||||
if [[ -z "${MASTER_ADDR:-}" ]]; then
|
||||
nodes=( $(scontrol show hostnames "${SLURM_JOB_NODELIST}") )
|
||||
export MASTER_ADDR="${nodes[0]}"
|
||||
fi
|
||||
|
||||
export NCCL_P2P_DISABLE="${NCCL_P2P_DISABLE:-1}"
|
||||
export TORCH_NCCL_ENABLE_MONITORING="${TORCH_NCCL_ENABLE_MONITORING:-0}"
|
||||
export NCCL_DEBUG_SUBSYS="${NCCL_DEBUG_SUBSYS:-INIT,NET}"
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="${WANDB_API_KEY:-7ff8b6e8356924f7a6dd51a0342dd1a422ea9352}"
|
||||
export WANDB_MODE="${WANDB_MODE:-online}"
|
||||
|
||||
set +u
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate mhuo-fv
|
||||
set -u
|
||||
export PYTHONPATH="/mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor:${PYTHONPATH:-}"
|
||||
|
||||
echo "=== Distillation Training (Slurm) ==="
|
||||
echo "Config: ${CONFIG}"
|
||||
echo "Num GPUs: ${NUM_GPUS}"
|
||||
echo "Num Nodes: ${NNODES}"
|
||||
echo "Master: ${MASTER_ADDR}:${MASTER_PORT}"
|
||||
echo "Extra args: ${EXTRA_ARGS[*]:-}"
|
||||
echo "====================================="
|
||||
|
||||
srun torchrun \
|
||||
--nnodes "${NNODES}" \
|
||||
--nproc_per_node "${NUM_GPUS}" \
|
||||
--rdzv_backend c10d \
|
||||
--rdzv_endpoint "${MASTER_ADDR}:${MASTER_PORT}" \
|
||||
--node_rank "${SLURM_PROCID}" \
|
||||
fastvideo/train/entrypoint/train.py \
|
||||
--config "${CONFIG}" \
|
||||
"${EXTRA_ARGS[@]}"
|
||||
@@ -6,7 +6,7 @@
|
||||
#
|
||||
# Examples:
|
||||
# bash examples/train/run_slurm.sh examples/train/configs/example.yaml 2
|
||||
# bash examples/train/run_slurm.sh examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml 4 \
|
||||
# bash examples/train/run_slurm.sh examples/train/configs/distill_wan2.1_t2v_1.3B_dmd2.yaml 4 \
|
||||
# --training.optimizer.learning_rate 1e-5
|
||||
# bash examples/train/run_slurm.sh examples/train/configs/example.yaml 8 \
|
||||
# --training.checkpoint.resume_from_checkpoint outputs/my_run/checkpoint-1000
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
# Training Scenarios
|
||||
|
||||
End-to-end multi-step training pipelines. Each scenario directory contains
|
||||
all configs, scripts, and data needed to run a complete workflow.
|
||||
|
||||
```
|
||||
scenario/
|
||||
└── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
|
||||
```
|
||||
|
||||
See the `usage.md` inside each scenario for step-by-step instructions.
|
||||
|
||||
For single-step configs, see `examples/train/configs/`.
|
||||
@@ -1,92 +0,0 @@
|
||||
"""Decode 'real' latents from a KD cache to videos.
|
||||
|
||||
Uses fastvideo's AutoencoderKLWan and its built-in denormalization
|
||||
(latent * std + mean, matching normalize_dit_input's inverse).
|
||||
|
||||
Usage:
|
||||
python examples/train/scenario/ode_init_self_forcing_wan_causal/decode_kd_cache.py \
|
||||
--cache_dir data/kd_test_cache_small \
|
||||
--out_dir data/kd_test_cache_small/video \
|
||||
--fps 16
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
|
||||
MODEL_ID = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
|
||||
|
||||
def denormalize(latents: torch.Tensor, vae: AutoencoderKLWan) -> torch.Tensor:
|
||||
"""Inverse of normalize_dit_input('wan', ...).
|
||||
|
||||
normalize_dit_input: normalized = (latent - mean) * (1/std)
|
||||
inverse: latent = normalized * std + mean
|
||||
"""
|
||||
cfg = vae.config
|
||||
mean = torch.tensor(cfg.latents_mean, dtype=latents.dtype,
|
||||
device=latents.device).view(1, -1, 1, 1, 1)
|
||||
std = torch.tensor(cfg.latents_std, dtype=latents.dtype,
|
||||
device=latents.device).view(1, -1, 1, 1, 1)
|
||||
return latents * std + mean
|
||||
|
||||
|
||||
def decode_real(pt_path: Path, vae: AutoencoderKLWan,
|
||||
device: torch.device) -> np.ndarray:
|
||||
"""Load one .pt cache file and decode its 'real' latent to uint8 [T,H,W,3]."""
|
||||
d = torch.load(pt_path, weights_only=True, map_location="cpu")
|
||||
real = d["real"].float() # [T, C, H, W] (normalized)
|
||||
|
||||
# [T, C, H, W] → [1, C, T, H, W]
|
||||
latent = real.permute(1, 0, 2, 3).unsqueeze(0).to(device)
|
||||
latent = denormalize(latent, vae)
|
||||
|
||||
with torch.no_grad():
|
||||
video = vae.decode(latent).sample # [1, 3, T_px, H_px, W_px]
|
||||
|
||||
video = video.squeeze(0).permute(1, 2, 3, 0) # [T, H, W, 3]
|
||||
video = ((video.clamp(-1, 1) + 1) / 2 * 255).cpu().float().numpy()
|
||||
return video.astype(np.uint8)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--cache_dir", default="data/kd_test_cache_small")
|
||||
parser.add_argument("--out_dir", default="data/kd_test_cache_small/video")
|
||||
parser.add_argument("--model_id", default=MODEL_ID)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--overwrite", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
samples_dir = Path(args.cache_dir) / "samples"
|
||||
out_dir = Path(args.out_dir)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
print(f"Loading VAE from {args.model_id} ...")
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
args.model_id, subfolder="vae", torch_dtype=torch.float32)
|
||||
vae.eval().to(device)
|
||||
|
||||
pts = sorted(samples_dir.glob("*.pt"))
|
||||
print(f"Found {len(pts)} samples in {samples_dir}")
|
||||
|
||||
for pt in pts:
|
||||
out_path = out_dir / (pt.stem + ".mp4")
|
||||
if out_path.exists() and not args.overwrite:
|
||||
print(f" skip {pt.name} (exists)")
|
||||
continue
|
||||
video_np = decode_real(pt, vae, device)
|
||||
imageio.mimsave(str(out_path), video_np, fps=args.fps)
|
||||
print(f" {pt.name} → {out_path.name} {video_np.shape}")
|
||||
|
||||
print(f"Done. Videos in {out_dir}/")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,91 +0,0 @@
|
||||
# KD Test: Causal Wan 1.3B ODE-init
|
||||
#
|
||||
# Smoke-test for KDCausalMethod on a tiny 16-sample dataset.
|
||||
# Uses 1.3B as teacher (instead of 14B) so everything fits on 1 GPU.
|
||||
#
|
||||
# Step 1 — Run KD training:
|
||||
# torchrun --nproc_per_node=1 -m fastvideo.train.entrypoint.train \
|
||||
# --config examples/train/scenario/ode_init_self_forcing_wan_causal/step1_kd.yaml
|
||||
#
|
||||
# Step 2 — Export ode_init checkpoint:
|
||||
# python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
# --role student \
|
||||
# --checkpoint_dir outputs/kdtest/kd_causal/latest \
|
||||
# --output_dir outputs/kdtest/kd_ode_init
|
||||
#
|
||||
# Step 3 — Run Self-Forcing with the exported ode_init:
|
||||
# torchrun --nproc_per_node=1 -m fastvideo.train.entrypoint.train \
|
||||
# --config examples/train/scenario/ode_init_self_forcing_wan_causal/step3_self_forcing.yaml
|
||||
#
|
||||
# Success criteria:
|
||||
# - Cache generates 16 .pt files under data/kd_test_cache/samples/
|
||||
# - kd_loss decreases steadily and approaches ~0 by step 200
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher: # using 1.3B for single-GPU test; swap to 14B for real runs
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.knowledge_distillation.kd.KDCausalMethod
|
||||
teacher_path_cache: data/kd_test_cache_small
|
||||
t_list: [995, 937, 833, 625, 0]
|
||||
student_sample_steps: 4
|
||||
teacher_inference_steps: 48
|
||||
teacher_guidance_scale: 3.5
|
||||
num_frames_per_block: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_small
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 20
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 7e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/kdtest/kd_causal
|
||||
training_state_checkpointing_steps: 300
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
project_name: kd-test
|
||||
run_name: kd_causal_smoke
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,21 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# Step 2 — Export ODE-init weights from KD checkpoint to diffusers format.
|
||||
#
|
||||
# Usage:
|
||||
# bash examples/train/scenario/ode_init_self_forcing_wan_causal/step2_export.sh
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
CHECKPOINT_DIR="${1:-outputs/kdtest/kd_causal/checkpoint-300}"
|
||||
OUTPUT_DIR="${2:-outputs/kdtest/kd_ode_init}"
|
||||
|
||||
echo "Exporting ODE-init weights..."
|
||||
echo " checkpoint: ${CHECKPOINT_DIR}"
|
||||
echo " output: ${OUTPUT_DIR}"
|
||||
|
||||
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
--role student \
|
||||
--checkpoint "${CHECKPOINT_DIR}" \
|
||||
--output-dir "${OUTPUT_DIR}"
|
||||
|
||||
echo "Done. Use the exported weights in step3_self_forcing.yaml."
|
||||
@@ -1,111 +0,0 @@
|
||||
# Self-Forcing Test: Causal Wan 1.3B
|
||||
#
|
||||
# Smoke-test for SelfForcingMethod on a tiny 16-sample dataset.
|
||||
# Run this AFTER step1_kd.yaml + dcp_to_diffusers export.
|
||||
#
|
||||
# Prerequisite — export ode_init from the KD run:
|
||||
# python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
# --role student \
|
||||
# --checkpoint_dir outputs/kdtest/kd_causal/latest \
|
||||
# --output_dir outputs/kdtest/kd_ode_init
|
||||
#
|
||||
# Then run this config:
|
||||
# torchrun --nproc_per_node=1 -m fastvideo.train.entrypoint.train \
|
||||
# --config examples/train/scenario/ode_init_self_forcing_wan_causal/step3_self_forcing.yaml
|
||||
#
|
||||
# Success criteria:
|
||||
# - generator_loss and critic_loss both finite at step 0
|
||||
# - Validation videos logged to tracker every 20 steps
|
||||
# - No crash through max_train_steps=100
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
transformer_override_safetensor: outputs/kdtest/kd_ode_init/transformer/model.safetensors
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 4.5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_small
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 20
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1e-5
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 100
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/kdtest/self_forcing_causal
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
project_name: kd-test
|
||||
run_name: self_forcing_causal_smoke
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
|
||||
dataset_file: examples/train/scenario/ode_init_self_forcing_wan_causal/validation.json
|
||||
every_steps: 20
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 69
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,19 +0,0 @@
|
||||
# ODE-Init Self-Forcing: Wan 2.1 Causal 1.3B
|
||||
|
||||
End-to-end scenario that trains a causal Wan 1.3B model via
|
||||
ODE-init knowledge distillation followed by self-forcing.
|
||||
|
||||
## Steps
|
||||
|
||||
```bash
|
||||
# Step 1 — KD training (produces ODE-init checkpoint)
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/ode_init_self_forcing_wan_causal/step1_kd.yaml
|
||||
|
||||
# Step 2 — Export ODE-init weights to diffusers format
|
||||
bash examples/train/scenario/ode_init_self_forcing_wan_causal/step2_export.sh
|
||||
|
||||
# Step 3 — Self-forcing with the exported ODE-init
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/ode_init_self_forcing_wan_causal/step3_self_forcing.yaml
|
||||
```
|
||||
@@ -1,20 +0,0 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights.",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 4,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 69
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The scene is captured from a ground-level angle, giving a low and intimate perspective.",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 4,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 69
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "00 Val-00: W",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "01 Val-01: S",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "02 Val-02: A",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "03 Val-03: D",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "00 Val-00: W",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "01 Val-01: S",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "02 Val-02: A",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "03 Val-03: D",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "04 Val-04: u",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/u.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "05 Val-05: d",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/d.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "06 Val-06: l",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/l.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "07 Val-07: r",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/r.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,324 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "00 Val-00: W",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "01 Val-01: S",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "02 Val-02: A",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "03 Val-03: D",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "04 Val-04: u",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/u.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "05 Val-05: d",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/d.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "06 Val-06: l",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/l.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "07 Val-07: r",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/r.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "08 Val-00: key rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "09 Val-01: key rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_2.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "10 Val-02: camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "11 Val-03: camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_2.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "12 Val-00: key+camera excl rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "13 Val-01: key+camera excl rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_2.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "14 Val-02: key+camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "15 Val-03: key+camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_1_action_rand_2.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "16 Val-04: (simultaneous) key rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_2_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "17 Val-05: (simultaneous) camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_2_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "18 Val-06: (simultaneous) key+camera excl rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_2_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "19 Val-07: (simultaneous) key+camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_2_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "20 Val-08: W+A",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/WA.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "21 Val-09: S+u",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000013.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S_u.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "22 Val-08: Still",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000005.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/still.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "23 Val-09: Still",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000013.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/still.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "24 Val-06: key+camera excl rand Frame 4",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1_f4.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "25 Val-07: key+camera excl rand Frame 4",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_2_f4.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "26 Train-00",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/first_frame/000500.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/videos/000500_action.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "27 Train-01",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/first_frame/001000.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/videos/001000_action.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "28 Doom-00: W",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000000.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "29 Doom-01: key rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000001.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "30 Doom-02: camera rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000002.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "31 Doom-03: key+camera excl rand",
|
||||
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000003.jpg",
|
||||
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1.npy",
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 352,
|
||||
"width": 640,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,7 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.attention.layer import (DistributedAttention, DistributedAttention_VSA, LocalAttention)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.attention.layer import (DistributedAttention,
|
||||
DistributedAttention_VSA, LocalAttention)
|
||||
from fastvideo.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -54,13 +54,17 @@ class AttentionMetadata:
|
||||
# Current step of diffusion process
|
||||
current_timestep: int
|
||||
|
||||
def asdict_zerocopy(self, skip_fields: set[str] | None = None) -> dict[str, Any]:
|
||||
def asdict_zerocopy(self,
|
||||
skip_fields: set[str] | None = None) -> dict[str, Any]:
|
||||
"""Similar to dataclasses.asdict, but avoids deepcopying."""
|
||||
if skip_fields is None:
|
||||
skip_fields = set()
|
||||
# Note that if we add dataclasses as fields, they will need
|
||||
# similar handling.
|
||||
return {field.name: getattr(self, field.name) for field in fields(self) if field.name not in skip_fields}
|
||||
return {
|
||||
field.name: getattr(self, field.name)
|
||||
for field in fields(self) if field.name not in skip_fields
|
||||
}
|
||||
|
||||
|
||||
T = TypeVar("T", bound=AttentionMetadata)
|
||||
@@ -121,7 +125,8 @@ class AttentionImpl(ABC, Generic[T]):
|
||||
) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def preprocess_qkv(self, qkv: torch.Tensor, attn_metadata: T) -> torch.Tensor:
|
||||
def preprocess_qkv(self, qkv: torch.Tensor,
|
||||
attn_metadata: T) -> torch.Tensor:
|
||||
"""Preprocess QKV tensor before performing attention operation.
|
||||
|
||||
Default implementation returns the tensor unchanged.
|
||||
|
||||
@@ -72,11 +72,12 @@ class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> FlashAttnMetadata:
|
||||
return FlashAttnMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
return FlashAttnMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class FlashAttentionImpl(AttentionImpl):
|
||||
@@ -102,24 +103,32 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
|
||||
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor, key_len: int) -> torch.Tensor:
|
||||
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor,
|
||||
key_len: int) -> torch.Tensor:
|
||||
# Normalize attn_mask to [B, key_len] where True means valid token.
|
||||
if attn_mask.dim() == 4:
|
||||
attn_mask = attn_mask[:, 0, 0, :]
|
||||
elif attn_mask.dim() == 3:
|
||||
attn_mask = attn_mask[:, 0, :]
|
||||
elif attn_mask.dim() != 2:
|
||||
raise ValueError(f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}")
|
||||
raise ValueError(
|
||||
f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}"
|
||||
)
|
||||
|
||||
# SDPA additive mask convention: valid=0, masked=-inf/large negative.
|
||||
key_padding_mask = attn_mask if attn_mask.dtype == torch.bool else attn_mask >= 0
|
||||
if attn_mask.dtype == torch.bool:
|
||||
key_padding_mask = attn_mask
|
||||
else:
|
||||
# SDPA additive mask convention: valid=0, masked=-inf/large negative.
|
||||
key_padding_mask = attn_mask >= 0
|
||||
|
||||
if key_padding_mask.shape[-1] != key_len:
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected {key_len}, got {key_padding_mask.shape[-1]}")
|
||||
raise ValueError(
|
||||
"Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected {key_len}, got {key_padding_mask.shape[-1]}")
|
||||
return key_padding_mask
|
||||
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask")
|
||||
and attn_metadata.attn_mask is not None):
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
@@ -135,7 +144,8 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
dtype=torch.bool,
|
||||
device=query.device,
|
||||
)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, key.shape[1]).to(device=key.device)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(
|
||||
attn_mask, key.shape[1]).to(device=key.device)
|
||||
return flash_attn_varlen_qk_no_pad(
|
||||
query,
|
||||
key,
|
||||
@@ -149,8 +159,13 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
|
||||
output = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0, softmax_scale=None)
|
||||
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0),
|
||||
value=True)
|
||||
output = flash_attn_no_pad(qkv,
|
||||
attn_mask,
|
||||
causal=False,
|
||||
dropout_p=0,
|
||||
softmax_scale=None)
|
||||
else:
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
|
||||
@@ -3,7 +3,9 @@
|
||||
import torch
|
||||
from sageattn3 import sageattn3_blackwell
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
|
||||
@@ -3,7 +3,8 @@
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -45,11 +46,12 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
return SDPAMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
@@ -80,8 +82,9 @@ class SDPAImpl(AttentionImpl):
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_mask = attn_metadata.attn_mask if (
|
||||
attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
@@ -90,6 +93,7 @@ class SDPAImpl(AttentionImpl):
|
||||
}
|
||||
if query.shape[1] != key.shape[1]:
|
||||
attn_kwargs["enable_gqa"] = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(query, key, value, **attn_kwargs)
|
||||
output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query, key, value, **attn_kwargs)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
@@ -55,11 +55,13 @@ def compress_kernel(
|
||||
|
||||
x_offset = idx_bh * L * D
|
||||
xm_offset = idx_bh * ((L + BLOCK_L - 1) // BLOCK_L) * D
|
||||
x = tl.load(X + x_offset + offs_l[:, None] * D + offs_d[None, :], mask=offs_l[:, None] < L)
|
||||
x = tl.load(X + x_offset + offs_l[:, None] * D + offs_d[None, :],
|
||||
mask=offs_l[:, None] < L)
|
||||
|
||||
nx = min(BLOCK_L, L - idx_l * BLOCK_L)
|
||||
x_mean = tl.sum(x, axis=0, dtype=tl.float32) / nx
|
||||
tl.store(XM + xm_offset + idx_l * D + offs_d, x_mean.to(XM.dtype.element_ty))
|
||||
tl.store(XM + xm_offset + idx_l * D + offs_d,
|
||||
x_mean.to(XM.dtype.element_ty))
|
||||
|
||||
|
||||
def mean_pool(x: torch.Tensor, BLK: int) -> torch.Tensor:
|
||||
@@ -96,7 +98,8 @@ def get_block_map(
|
||||
lut: Top-k indices of shape (B, H, num_q_blocks, topk)
|
||||
topk: Number of key blocks selected
|
||||
"""
|
||||
arg_k = k - torch.mean(k, dim=-2, keepdim=True) # smooth-k technique from SageAttention
|
||||
arg_k = k - torch.mean(
|
||||
k, dim=-2, keepdim=True) # smooth-k technique from SageAttention
|
||||
pooled_qblocks = mean_pool(q, BLKQ)
|
||||
pooled_kblocks = mean_pool(arg_k, BLKK)
|
||||
pooled_score = pooled_qblocks @ pooled_kblocks.transpose(-1, -2)
|
||||
@@ -297,7 +300,11 @@ class SLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q, k, topk_ratio=topk_ratio, BLKQ=self.BLKQ, BLKK=self.BLKK)
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=self.BLKQ,
|
||||
BLKK=self.BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
@@ -305,7 +312,8 @@ class SLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
v = v.to(self.dtype)
|
||||
|
||||
# Sparse attention
|
||||
o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, self.BLKQ, self.BLKK)
|
||||
o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, self.BLKQ,
|
||||
self.BLKK)
|
||||
|
||||
# Linear attention with feature maps
|
||||
q_linear = self.feature_map_q(q).contiguous().to(self.dtype)
|
||||
@@ -404,10 +412,14 @@ class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
nn.Module.__init__(self)
|
||||
|
||||
if not SAGESLA_ENABLED:
|
||||
raise ImportError("SageSLA requires spas_sage_attn. "
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
raise ImportError(
|
||||
"SageSLA requires spas_sage_attn. "
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git"
|
||||
)
|
||||
|
||||
assert head_size in [64, 128], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
assert head_size in [
|
||||
64, 128
|
||||
], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
@@ -497,7 +509,11 @@ class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
BLKQ, BLKK = 128, 64
|
||||
|
||||
# Compute block-sparse attention pattern
|
||||
sparse_map, lut, real_topk = get_block_map(q, k, topk_ratio=topk_ratio, BLKQ=BLKQ, BLKK=BLKK)
|
||||
sparse_map, lut, real_topk = get_block_map(q,
|
||||
k,
|
||||
topk_ratio=topk_ratio,
|
||||
BLKQ=BLKQ,
|
||||
BLKK=BLKK)
|
||||
|
||||
# Convert to compute dtype
|
||||
q = q.to(self.dtype)
|
||||
@@ -510,33 +526,47 @@ class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
scale = 1.0 / (headdim**0.5)
|
||||
|
||||
# Quantize Q, K to INT8
|
||||
q_int8, q_scale, k_int8, k_scale = get_vanilla_qk_quant(q, k, km, BLKQ, BLKK)
|
||||
q_int8, q_scale, k_int8, k_scale = get_vanilla_qk_quant(
|
||||
q, k, km, BLKQ, BLKK)
|
||||
lut_triton, valid_block_num = block_map_lut_triton(sparse_map)
|
||||
|
||||
# Quantize V to FP8
|
||||
b, h_kv, kv_len, head_dim = v.shape
|
||||
padded_len = (kv_len + 127) // 128 * 128
|
||||
v_transposed_permutted = torch.empty((b, h_kv, head_dim, padded_len), dtype=v.dtype, device=v.device)
|
||||
v_transposed_permutted = torch.empty((b, h_kv, head_dim, padded_len),
|
||||
dtype=v.dtype,
|
||||
device=v.device)
|
||||
fused.transpose_pad_permute_cuda(v, v_transposed_permutted, 1)
|
||||
v_fp8 = torch.empty(v_transposed_permutted.shape, dtype=torch.float8_e4m3fn, device=v.device)
|
||||
v_scale = torch.empty((b, h_kv, head_dim), dtype=torch.float32, device=v.device)
|
||||
fused.scale_fuse_quant_cuda(v_transposed_permutted, v_fp8, v_scale, kv_len, 2.25, 1)
|
||||
v_fp8 = torch.empty(v_transposed_permutted.shape,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=v.device)
|
||||
v_scale = torch.empty((b, h_kv, head_dim),
|
||||
dtype=torch.float32,
|
||||
device=v.device)
|
||||
fused.scale_fuse_quant_cuda(v_transposed_permutted, v_fp8, v_scale,
|
||||
kv_len, 2.25, 1)
|
||||
|
||||
# Sparse attention with quantized kernels
|
||||
o_s = torch.empty_like(q)
|
||||
if arch == "sm90":
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_sm90(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, q_scale, k_scale, v_scale, 1, False, 1, scale)
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
q_scale, k_scale, v_scale, 1, False, 1, scale)
|
||||
else:
|
||||
pvthreshold = torch.full((q.shape[-3], ), 1e6, dtype=torch.float32, device=q.device)
|
||||
pvthreshold = torch.full((q.shape[-3], ),
|
||||
1e6,
|
||||
dtype=torch.float32,
|
||||
device=q.device)
|
||||
if SAGE2PP_ENABLED:
|
||||
qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, pvthreshold, q_scale, k_scale, v_scale, 1,
|
||||
False, 1, scale, 0)
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
else:
|
||||
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, pvthreshold, q_scale, k_scale, v_scale, 1,
|
||||
False, 1, scale, 0)
|
||||
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num,
|
||||
pvthreshold, q_scale, k_scale, v_scale, 1, False, 1, scale,
|
||||
0)
|
||||
# ========== END SPARGE ==========
|
||||
|
||||
# Linear attention with feature maps
|
||||
|
||||
@@ -12,7 +12,9 @@ except ImportError:
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -29,12 +31,14 @@ def get_tile_partition_indices(
|
||||
) -> torch.LongTensor:
|
||||
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)
|
||||
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),
|
||||
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())
|
||||
index = torch.cat(ls, dim=0)
|
||||
return index
|
||||
@@ -46,7 +50,8 @@ def get_reverse_tile_partition_indices(
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
|
||||
return torch.argsort(
|
||||
get_tile_partition_indices(dit_seq_shape, tile_size, device))
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
@@ -99,8 +104,10 @@ def get_non_pad_index(
|
||||
n_win = variable_block_sizes.shape[0]
|
||||
device = variable_block_sizes.device
|
||||
starts_pad = torch.arange(n_win, device=device) * max_block_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(max_block_size, device=device)[None, :]
|
||||
index_mask = torch.arange(max_block_size, device=device)[None, :] < variable_block_sizes[:, None]
|
||||
index_pad = starts_pad[:, None] + torch.arange(max_block_size,
|
||||
device=device)[None, :]
|
||||
index_mask = torch.arange(
|
||||
max_block_size, device=device)[None, :] < variable_block_sizes[:, None]
|
||||
return index_pad[index_mask]
|
||||
|
||||
|
||||
@@ -160,17 +167,23 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
**kwargs: dict[str, Any],
|
||||
) -> VideoSparseAttentionMetadata:
|
||||
patch_size = patch_size
|
||||
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
|
||||
dit_seq_shape = (raw_latent_shape[0] // patch_size[0],
|
||||
raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2])
|
||||
|
||||
num_tiles = (math.ceil(dit_seq_shape[0] / VSA_TILE_SIZE[0]), math.ceil(dit_seq_shape[1] / VSA_TILE_SIZE[1]),
|
||||
num_tiles = (math.ceil(dit_seq_shape[0] / VSA_TILE_SIZE[0]),
|
||||
math.ceil(dit_seq_shape[1] / VSA_TILE_SIZE[1]),
|
||||
math.ceil(dit_seq_shape[2] / VSA_TILE_SIZE[2]))
|
||||
total_seq_length = math.prod(dit_seq_shape)
|
||||
|
||||
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, VSA_TILE_SIZE, device)
|
||||
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, VSA_TILE_SIZE, device)
|
||||
variable_block_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device)
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, math.prod(VSA_TILE_SIZE))
|
||||
tile_partition_indices = get_tile_partition_indices(
|
||||
dit_seq_shape, VSA_TILE_SIZE, device)
|
||||
reverse_tile_partition_indices = get_reverse_tile_partition_indices(
|
||||
dit_seq_shape, VSA_TILE_SIZE, device)
|
||||
variable_block_sizes = construct_variable_block_sizes(
|
||||
dit_seq_shape, num_tiles, device)
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes,
|
||||
math.prod(VSA_TILE_SIZE))
|
||||
|
||||
return VideoSparseAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -200,19 +213,23 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
|
||||
def tile(self, x: torch.Tensor, num_tiles: list[int], tile_partition_indices: torch.LongTensor,
|
||||
def tile(self, x: torch.Tensor, num_tiles: list[int],
|
||||
tile_partition_indices: torch.LongTensor,
|
||||
non_pad_index: torch.LongTensor) -> torch.Tensor:
|
||||
t_padded_size = num_tiles[0] * VSA_TILE_SIZE[0]
|
||||
h_padded_size = num_tiles[1] * VSA_TILE_SIZE[1]
|
||||
w_padded_size = num_tiles[2] * VSA_TILE_SIZE[2]
|
||||
|
||||
x_padded = torch.zeros((x.shape[0], t_padded_size * h_padded_size * w_padded_size, x.shape[-2], x.shape[-1]),
|
||||
device=x.device,
|
||||
dtype=x.dtype)
|
||||
x_padded = torch.zeros(
|
||||
(x.shape[0], t_padded_size * h_padded_size * w_padded_size,
|
||||
x.shape[-2], x.shape[-1]),
|
||||
device=x.device,
|
||||
dtype=x.dtype)
|
||||
x_padded[:, non_pad_index] = x[:, tile_partition_indices]
|
||||
return x_padded
|
||||
|
||||
def untile(self, x: torch.Tensor, reverse_tile_partition_indices: torch.LongTensor,
|
||||
def untile(self, x: torch.Tensor,
|
||||
reverse_tile_partition_indices: torch.LongTensor,
|
||||
non_pad_index: torch.LongTensor) -> torch.Tensor:
|
||||
x = x[:, non_pad_index][:, reverse_tile_partition_indices]
|
||||
return x
|
||||
@@ -222,7 +239,8 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: VideoSparseAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return self.tile(qkv, attn_metadata.num_tiles, attn_metadata.tile_partition_indices,
|
||||
return self.tile(qkv, attn_metadata.num_tiles,
|
||||
attn_metadata.tile_partition_indices,
|
||||
attn_metadata.non_pad_index)
|
||||
|
||||
def postprocess_output(
|
||||
@@ -230,7 +248,8 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
output: torch.Tensor,
|
||||
attn_metadata: VideoSparseAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return self.untile(output, attn_metadata.reverse_tile_partition_indices, attn_metadata.non_pad_index)
|
||||
return self.untile(output, attn_metadata.reverse_tile_partition_indices,
|
||||
attn_metadata.non_pad_index)
|
||||
|
||||
def forward( # type: ignore[override]
|
||||
self,
|
||||
@@ -247,17 +266,20 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
|
||||
VSA_sparsity = attn_metadata.VSA_sparsity
|
||||
|
||||
cur_topk = math.ceil((1 - VSA_sparsity) * (attn_metadata.total_seq_length / math.prod(VSA_TILE_SIZE)))
|
||||
cur_topk = math.ceil(
|
||||
(1 - VSA_sparsity) *
|
||||
(attn_metadata.total_seq_length / math.prod(VSA_TILE_SIZE)))
|
||||
|
||||
if video_sparse_attn is None:
|
||||
raise NotImplementedError("video_sparse_attn is not installed")
|
||||
hidden_states = video_sparse_attn(query,
|
||||
key,
|
||||
value,
|
||||
attn_metadata.variable_block_sizes,
|
||||
attn_metadata.variable_block_sizes,
|
||||
cur_topk,
|
||||
block_size=VSA_TILE_SIZE,
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
hidden_states = video_sparse_attn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_metadata.variable_block_sizes,
|
||||
attn_metadata.variable_block_sizes,
|
||||
cur_topk,
|
||||
block_size=VSA_TILE_SIZE,
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -6,8 +6,11 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo_kernel import (moba_attn_varlen, process_moba_input, process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
@@ -91,10 +94,12 @@ class VideoMobaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
) -> VideoMobaAttentionMetadata:
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[1] % patch_size[
|
||||
1] == 0 and raw_latent_shape[2] % patch_size[
|
||||
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[
|
||||
1] % patch_size[1] == 0 and raw_latent_shape[2] % patch_size[
|
||||
2] == 0, f"spatial patch_resolution {raw_latent_shape} should be divisible by patch_size {patch_size}"
|
||||
patch_resolution = [t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)]
|
||||
patch_resolution = [
|
||||
t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)
|
||||
]
|
||||
|
||||
return VideoMobaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -165,11 +170,19 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
moba_chunk_size = attn_metadata.st_chunk_size
|
||||
moba_topk = attn_metadata.st_topk
|
||||
|
||||
query, chunk_size = process_moba_input(query, attn_metadata.patch_resolution, moba_chunk_size)
|
||||
key, chunk_size = process_moba_input(key, attn_metadata.patch_resolution, moba_chunk_size)
|
||||
value, chunk_size = process_moba_input(value, attn_metadata.patch_resolution, moba_chunk_size)
|
||||
query, chunk_size = process_moba_input(query,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
key, chunk_size = process_moba_input(key,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
value, chunk_size = process_moba_input(value,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
max_seqlen = query.shape[1]
|
||||
indices_q = torch.arange(0, query.shape[0] * query.shape[1], device=query.device)
|
||||
indices_q = torch.arange(0,
|
||||
query.shape[0] * query.shape[1],
|
||||
device=query.device)
|
||||
cu_seqlens = torch.arange(0,
|
||||
query.shape[0] * query.shape[1] + 1,
|
||||
query.shape[1],
|
||||
@@ -192,7 +205,10 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
simsum_threshold=attn_metadata.moba_threshold,
|
||||
threshold_type=attn_metadata.moba_threshold_type,
|
||||
)
|
||||
hidden_states = self.pad_input(hidden_states, indices_q, batch_size, sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states, attn_metadata.patch_resolution, moba_chunk_size)
|
||||
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -4,9 +4,10 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.attention.selector import backend_name_to_enum, get_attn_backend
|
||||
from fastvideo.distributed.communication_op import (sequence_model_parallel_all_gather,
|
||||
sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.distributed.parallel_state import (get_sp_parallel_rank, get_sp_world_size)
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
@@ -37,7 +38,10 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
@@ -85,7 +89,8 @@ class DistributedAttention(nn.Module):
|
||||
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
batch_size, _, num_heads, _ = q.shape
|
||||
local_rank = get_sp_parallel_rank()
|
||||
world_size = get_sp_world_size()
|
||||
@@ -94,10 +99,13 @@ class DistributedAttention(nn.Module):
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3*batch, seq_len, num_heads, head_dim]
|
||||
qkv = torch.cat([q, k, v],
|
||||
dim=0) # [3*batch, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# After all-to-all, each rank has the full sequence but only a subset of heads.
|
||||
# Trim away SP padding for attention compute, then pad back before returning.
|
||||
@@ -107,17 +115,23 @@ class DistributedAttention(nn.Module):
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
qkv[:batch_size * 2] = _apply_rotary_emb(qkv[:batch_size * 2], cos, sin, is_neox_style=False)
|
||||
qkv[:batch_size * 2] = _apply_rotary_emb(qkv[:batch_size * 2],
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
assert replicated_k is not None and replicated_v is not None
|
||||
replicated_qkv = torch.cat([replicated_q, replicated_k, replicated_v],
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
replicated_qkv = torch.cat(
|
||||
[replicated_q, replicated_k, replicated_v],
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
heads_per_rank = num_heads // world_size
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank * heads_per_rank:(local_rank + 1) * heads_per_rank]
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank *
|
||||
heads_per_rank:(local_rank + 1) *
|
||||
heads_per_rank]
|
||||
qkv = torch.cat([qkv, replicated_qkv], dim=1)
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
@@ -131,13 +145,16 @@ class DistributedAttention(nn.Module):
|
||||
replicated_output = output[:, split_idx:]
|
||||
output = output[:, :split_idx]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(replicated_output.contiguous(), dim=2)
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_seq_len))
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
|
||||
return output, replicated_output
|
||||
|
||||
@@ -179,19 +196,23 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
# Check text tokens are not supported for VSA now
|
||||
assert replicated_q is None and replicated_k is None and replicated_v is None, "Replicated QKV is not supported for VSA now"
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
# Stack QKV
|
||||
qkvg = torch.cat([q, k, v, gate_compress], dim=0) # [4*batch, seq_len, num_heads, head_dim]
|
||||
qkvg = torch.cat([q, k, v, gate_compress],
|
||||
dim=0) # [4*batch, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
# Before: [4*batch, shard_seq_len, num_heads, head_dim]
|
||||
# After: [4*batch, full_seq_len, shard_num_heads, head_dim]
|
||||
qkvg = sequence_model_parallel_all_to_all_4D(qkvg, scatter_dim=2, gather_dim=1)
|
||||
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# After all-to-all, each rank has the full sequence but only a subset of heads
|
||||
pad_seq_len = qkvg.shape[1] - original_seq_len
|
||||
@@ -199,12 +220,16 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
qkvg[:batch_size * 2] = _apply_rotary_emb(qkvg[:batch_size * 2], cos, sin, is_neox_style=False)
|
||||
qkvg[:batch_size * 2] = _apply_rotary_emb(qkvg[:batch_size * 2],
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
|
||||
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
|
||||
|
||||
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
|
||||
output = self.attn_impl.forward(q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
|
||||
output = self.attn_impl.forward(
|
||||
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
@@ -214,7 +239,9 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
|
||||
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_seq_len))
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
@@ -240,7 +267,10 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
@@ -273,7 +303,8 @@ class LocalAttention(nn.Module):
|
||||
torch.Tensor: Output tensor after local attention
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
@@ -43,7 +43,8 @@ def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
|
||||
* None otherwise
|
||||
'''
|
||||
backend_name = os.environ.get(STR_BACKEND_ENV_VAR)
|
||||
return (None if backend_name is None else backend_name_to_enum(backend_name))
|
||||
return (None
|
||||
if backend_name is None else backend_name_to_enum(backend_name))
|
||||
|
||||
|
||||
# Global state allows a particular choice of backend
|
||||
@@ -56,7 +57,8 @@ def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
|
||||
forced_attn_backend: AttentionBackendEnum | None = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None:
|
||||
def global_force_attn_backend(
|
||||
attn_backend: AttentionBackendEnum | None) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -85,7 +87,8 @@ def get_attn_backend(
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
|
||||
|
||||
@cache
|
||||
@@ -103,7 +106,8 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: AttentionBackendEnum | None = (get_global_forced_attn_backend())
|
||||
backend_by_global_setting: AttentionBackendEnum | None = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
else:
|
||||
@@ -117,14 +121,17 @@ def _cached_get_attn_backend(
|
||||
|
||||
if selected_backend not in supported_attention_backends:
|
||||
selected_backend = None
|
||||
attention_cls = current_platform.get_attn_backend_cls(selected_backend, head_size, dtype)
|
||||
attention_cls = current_platform.get_attn_backend_cls(
|
||||
selected_backend, head_size, dtype)
|
||||
if not attention_cls:
|
||||
raise ValueError(f"Invalid attention backend for {current_platform.device_name}")
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {current_platform.device_name}")
|
||||
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
|
||||
@@ -5,12 +5,16 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
raise ImportError(
|
||||
"flash_attn.cute is only available on CUDA devices; this error must be handled internally"
|
||||
)
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
raise NotImplementedError(
|
||||
f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})"
|
||||
)
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
@@ -57,7 +61,8 @@ def _flash_attn_cute_forward_fake(
|
||||
return out, lse
|
||||
|
||||
|
||||
def _flash_attn_cute_setup_context(ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
|
||||
def _flash_attn_cute_setup_context(ctx: torch.autograd.function.FunctionCtx,
|
||||
inputs, output) -> None:
|
||||
q, k, v, softmax_scale, causal, deterministic = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(q, k, v, out, lse)
|
||||
@@ -155,7 +160,8 @@ def _flash_attn_cute_varlen_forward_fake(
|
||||
return out, lse
|
||||
|
||||
|
||||
def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
|
||||
def _flash_attn_cute_varlen_setup_context(
|
||||
ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
|
||||
(
|
||||
q,
|
||||
k,
|
||||
@@ -223,7 +229,8 @@ def flash_attn_func(
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(
|
||||
q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
|
||||
|
||||
|
||||
@@ -42,9 +42,13 @@ def flash_attn_no_pad(
|
||||
seqlen = qkv.shape[1]
|
||||
nheads = qkv.shape[-2]
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
|
||||
x, key_padding_mask)
|
||||
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
|
||||
x_unpad = rearrange(x_unpad,
|
||||
"nnz (three h d) -> nnz three h d",
|
||||
three=3,
|
||||
h=nheads)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad,
|
||||
cu_seqlens,
|
||||
@@ -84,10 +88,12 @@ def flash_attn_no_pad_v3(
|
||||
batch_size, seqlen, _, nheads, head_dim = qkv.shape
|
||||
query, key, value = qkv.unbind(dim=2)
|
||||
|
||||
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
|
||||
rearrange(query, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(
|
||||
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(
|
||||
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
|
||||
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
@@ -132,10 +138,12 @@ def flash_attn_varlen_qk_no_pad(
|
||||
):
|
||||
batch_size, q_seqlen, nheads, _ = query.shape
|
||||
|
||||
query_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
query_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
|
||||
rearrange(query, "b s h d -> b s (h d)"), query_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(
|
||||
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(
|
||||
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
|
||||
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
|
||||
@@ -23,7 +23,8 @@ class DatasetType(str, Enum):
|
||||
return cls(value.lower())
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid dataset type: {value}. Must be one of: {', '.join([m.value for m in cls])}") from None
|
||||
f"Invalid dataset type: {value}. Must be one of: {', '.join([m.value for m in cls])}"
|
||||
) from None
|
||||
|
||||
@classmethod
|
||||
def choices(cls) -> list[str]:
|
||||
@@ -45,7 +46,8 @@ class VideoLoaderType(str, Enum):
|
||||
return cls(value.lower())
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid video loader: {value}. Must be one of: {', '.join([m.value for m in cls])}") from None
|
||||
f"Invalid video loader: {value}. Must be one of: {', '.join([m.value for m in cls])}"
|
||||
) from None
|
||||
|
||||
@classmethod
|
||||
def choices(cls) -> list[str]:
|
||||
@@ -90,7 +92,8 @@ class PreprocessConfig:
|
||||
seed: int = 42
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser, prefix: str = "preprocess") -> FlexibleArgumentParser:
|
||||
def add_cli_args(parser: FlexibleArgumentParser,
|
||||
prefix: str = "preprocess") -> FlexibleArgumentParser:
|
||||
"""Add preprocessing configuration arguments to the parser."""
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
|
||||
@@ -100,19 +103,22 @@ class PreprocessConfig:
|
||||
type=str,
|
||||
default=PreprocessConfig.model_path,
|
||||
help="Path to the model for preprocessing")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-path",
|
||||
type=str,
|
||||
default=PreprocessConfig.dataset_path,
|
||||
help="Path to the dataset directory for preprocessing")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-type",
|
||||
type=str,
|
||||
choices=DatasetType.choices(),
|
||||
default=PreprocessConfig.dataset_type.value,
|
||||
help="Type of the dataset")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-output-dir",
|
||||
type=str,
|
||||
default=PreprocessConfig.dataset_output_dir,
|
||||
help="The output directory where the dataset will be written.")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}dataset-path",
|
||||
type=str,
|
||||
default=PreprocessConfig.dataset_path,
|
||||
help="Path to the dataset directory for preprocessing")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}dataset-type",
|
||||
type=str,
|
||||
choices=DatasetType.choices(),
|
||||
default=PreprocessConfig.dataset_type.value,
|
||||
help="Type of the dataset")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}dataset-output-dir",
|
||||
type=str,
|
||||
default=PreprocessConfig.dataset_output_dir,
|
||||
help="The output directory where the dataset will be written.")
|
||||
|
||||
# Dataloader
|
||||
preprocess_args.add_argument(
|
||||
@@ -120,11 +126,13 @@ class PreprocessConfig:
|
||||
type=int,
|
||||
default=PreprocessConfig.dataloader_num_workers,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}preprocess-video-batch-size",
|
||||
type=int,
|
||||
default=PreprocessConfig.preprocess_video_batch_size,
|
||||
help="Batch size (per device) for the training dataloader.")
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process."
|
||||
)
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}preprocess-video-batch-size",
|
||||
type=int,
|
||||
default=PreprocessConfig.preprocess_video_batch_size,
|
||||
help="Batch size (per device) for the training dataloader.")
|
||||
|
||||
# Saver
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}samples-per-file",
|
||||
@@ -137,11 +145,12 @@ class PreprocessConfig:
|
||||
help="How often to save to parquet files")
|
||||
|
||||
# Video processing parameters
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}video-loader-type",
|
||||
type=str,
|
||||
choices=VideoLoaderType.choices(),
|
||||
default=PreprocessConfig.video_loader_type.value,
|
||||
help="Type of the video loader")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}video-loader-type",
|
||||
type=str,
|
||||
choices=VideoLoaderType.choices(),
|
||||
default=PreprocessConfig.video_loader_type.value,
|
||||
help="Type of the video loader")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}max-height",
|
||||
type=int,
|
||||
default=PreprocessConfig.max_height,
|
||||
@@ -154,10 +163,11 @@ class PreprocessConfig:
|
||||
type=int,
|
||||
default=PreprocessConfig.num_frames,
|
||||
help="Number of frames to process")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}video-length-tolerance-range",
|
||||
type=float,
|
||||
default=PreprocessConfig.video_length_tolerance_range,
|
||||
help="Video length tolerance range")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}video-length-tolerance-range",
|
||||
type=float,
|
||||
default=PreprocessConfig.video_length_tolerance_range,
|
||||
help="Video length tolerance range")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}train-fps",
|
||||
type=int,
|
||||
default=PreprocessConfig.train_fps,
|
||||
@@ -170,10 +180,11 @@ class PreprocessConfig:
|
||||
type=float,
|
||||
default=PreprocessConfig.drop_short_ratio,
|
||||
help="Ratio for dropping short videos")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}do-temporal-sample",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.do_temporal_sample,
|
||||
help="Whether to do temporal sampling")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}do-temporal-sample",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.do_temporal_sample,
|
||||
help="Whether to do temporal sampling")
|
||||
|
||||
# Model Training configuration
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
|
||||
@@ -192,15 +203,20 @@ class PreprocessConfig:
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls, kwargs: dict[str, Any]) -> Optional["PreprocessConfig"]:
|
||||
def from_kwargs(cls, kwargs: dict[str,
|
||||
Any]) -> Optional["PreprocessConfig"]:
|
||||
"""Create PreprocessConfig from keyword arguments."""
|
||||
if 'dataset_type' in kwargs and isinstance(kwargs['dataset_type'], str):
|
||||
kwargs['dataset_type'] = DatasetType.from_string(kwargs['dataset_type'])
|
||||
if 'video_loader_type' in kwargs and isinstance(kwargs['video_loader_type'], str):
|
||||
kwargs['video_loader_type'] = VideoLoaderType.from_string(kwargs['video_loader_type'])
|
||||
kwargs['dataset_type'] = DatasetType.from_string(
|
||||
kwargs['dataset_type'])
|
||||
if 'video_loader_type' in kwargs and isinstance(
|
||||
kwargs['video_loader_type'], str):
|
||||
kwargs['video_loader_type'] = VideoLoaderType.from_string(
|
||||
kwargs['video_loader_type'])
|
||||
|
||||
preprocess_config = cls()
|
||||
if not update_config_from_args(preprocess_config, kwargs, prefix="preprocess", pop_args=True):
|
||||
if not update_config_from_args(
|
||||
preprocess_config, kwargs, prefix="preprocess", pop_args=True):
|
||||
return None
|
||||
return preprocess_config
|
||||
|
||||
|
||||
@@ -2,7 +2,9 @@ from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig, LTX2AudioEncoderConfig, LTX2VocoderConfig)
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig)
|
||||
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -15,14 +15,17 @@ class LTX2AudioArchConfig(ArchConfig):
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioEncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2AudioEncoder"]))
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioEncoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioDecoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2AudioDecoder"]))
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioDecoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VocoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2Vocoder"]))
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2Vocoder"]))
|
||||
|
||||
@@ -13,7 +13,8 @@ logger = init_logger(__name__)
|
||||
@dataclass
|
||||
class ArchConfig:
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(
|
||||
default_factory=list) # mapping from huggingface weight names to custom names
|
||||
default_factory=list
|
||||
) # mapping from huggingface weight names to custom names
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -28,7 +29,8 @@ class ModelConfig:
|
||||
# Only called if 'name' is not found in ModelConfig directly
|
||||
if hasattr(self.arch_config, name):
|
||||
return getattr(self.arch_config, name)
|
||||
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
|
||||
raise AttributeError(
|
||||
f"'{type(self).__name__}' object has no attribute '{name}'")
|
||||
|
||||
def __getstate__(self):
|
||||
# Return a dictionary of attributes to pickle
|
||||
@@ -49,7 +51,8 @@ class ModelConfig:
|
||||
if key in valid_fields:
|
||||
setattr(arch_config, key, value)
|
||||
else:
|
||||
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
|
||||
raise AttributeError(
|
||||
f"{type(arch_config).__name__} has no field '{key}'")
|
||||
|
||||
if hasattr(arch_config, "__post_init__"):
|
||||
arch_config.__post_init__()
|
||||
@@ -63,7 +66,8 @@ class ModelConfig:
|
||||
if key in valid_fields:
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
logger.warning("%s does not contain field '%s'!", type(self).__name__, key)
|
||||
logger.warning("%s does not contain field '%s'!",
|
||||
type(self).__name__, key)
|
||||
raise AttributeError(f"Invalid field: {key}")
|
||||
|
||||
if hasattr(self, "__post_init__"):
|
||||
|
||||
@@ -7,9 +7,11 @@ from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig",
|
||||
"WanVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig",
|
||||
"WanGameVideoConfig"
|
||||
]
|
||||
|
||||
@@ -14,12 +14,11 @@ class DiTArchConfig(ArchConfig):
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum,
|
||||
...] = (AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_THREE,
|
||||
AttentionBackendEnum.SLA_ATTN, AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_THREE,
|
||||
AttentionBackendEnum.SLA_ATTN, AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -10,7 +10,8 @@ def is_transformer_blocks(n: str, m) -> bool:
|
||||
|
||||
@dataclass
|
||||
class CosmosArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
@@ -18,35 +19,58 @@ class CosmosArchConfig(DiTArchConfig):
|
||||
r"^time_embed\.time_proj\.(.*)$": r"time_embed.time_proj.\1",
|
||||
r"^time_embed\.t_embedder\.(.*)$": r"time_embed.t_embedder.\1",
|
||||
r"^time_embed\.norm\.(.*)$": r"time_embed.norm.\1",
|
||||
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\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"transformer_blocks.\1.attn1.norm_k.\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\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_q\.(.*)$": r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_k\.(.*)$": r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
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\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.\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\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
r"^norm_out\.(.*)$": r"norm_out.\1",
|
||||
r"^proj_out\.(.*)$": r"proj_out.\1",
|
||||
})
|
||||
|
||||
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+)\.ff\.(.*)$": r"transformer_blocks.\1.ff.\2",
|
||||
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+)\.ff\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.\2",
|
||||
})
|
||||
|
||||
# Cosmos-specific config parameters based on transformer_cosmos.py
|
||||
|
||||
@@ -12,33 +12,46 @@ def is_transformer_blocks(n: str, m) -> bool:
|
||||
class Cosmos25ArchConfig(DiTArchConfig):
|
||||
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Remove "net." prefix and map official structure to FastVideo
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$":
|
||||
r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
|
||||
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",
|
||||
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",
|
||||
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
|
||||
# Note: This also handles _extra_state if present
|
||||
r"^net\.t_embedding_norm\.(.*)$": r"time_embed.norm.\1",
|
||||
r"^net\.t_embedding_norm\.(.*)$":
|
||||
r"time_embed.norm.\1",
|
||||
|
||||
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
|
||||
r"^net\.crossattn_proj\.0\.weight$": r"crossattn_proj.0.weight",
|
||||
r"^net\.crossattn_proj\.0\.bias$": r"crossattn_proj.0.bias",
|
||||
r"^net\.crossattn_proj\.0\.weight$":
|
||||
r"crossattn_proj.0.weight",
|
||||
r"^net\.crossattn_proj\.0\.bias$":
|
||||
r"crossattn_proj.0.bias",
|
||||
|
||||
# Transformer blocks: net.blocks.N -> transformer_blocks.N
|
||||
# Self-attention (self_attn -> attn1)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$": r"transformer_blocks.\1.attn1.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$": r"transformer_blocks.\1.attn1.norm_k.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.weight",
|
||||
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_q._extra_state",
|
||||
@@ -46,12 +59,18 @@ class Cosmos25ArchConfig(DiTArchConfig):
|
||||
r"transformer_blocks.\1.attn1.norm_k._extra_state",
|
||||
|
||||
# Cross-attention (cross_attn -> attn2)
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$": r"transformer_blocks.\1.attn2.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$": r"transformer_blocks.\1.attn2.norm_k.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.weight",
|
||||
# RMSNorm _extra_state keys for cross-attention
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_q._extra_state",
|
||||
@@ -59,8 +78,10 @@ class Cosmos25ArchConfig(DiTArchConfig):
|
||||
r"transformer_blocks.\1.attn2.norm_k._extra_state",
|
||||
|
||||
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
|
||||
# These are now at the block level, not inside norm layers
|
||||
@@ -72,21 +93,27 @@ class Cosmos25ArchConfig(DiTArchConfig):
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$": r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$": r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
|
||||
|
||||
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm1.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm2.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$": r"transformer_blocks.\1.norm3.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$":
|
||||
r"transformer_blocks.\1.norm3.norm._extra_state",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
|
||||
r"^net\.final_layer\.linear\.(.*)$":
|
||||
r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
|
||||
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$": r"final_layer.linear_1.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$": r"final_layer.linear_2.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$":
|
||||
r"final_layer.linear_1.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$":
|
||||
r"final_layer.linear_2.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
|
||||
@@ -96,15 +123,24 @@ class Cosmos25ArchConfig(DiTArchConfig):
|
||||
|
||||
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",
|
||||
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",
|
||||
})
|
||||
|
||||
# Cosmos 2.5 specific config parameters
|
||||
|
||||
@@ -45,45 +45,63 @@ class HunyuanGameCraftArchConfig(DiTArchConfig):
|
||||
camera_net: bool = True
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_single_block, is_refiner_block, is_camera_net])
|
||||
default_factory=lambda:
|
||||
[is_double_block, is_single_block, is_refiner_block, is_camera_net])
|
||||
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
|
||||
|
||||
# Parameter names mapping from official checkpoint to FastVideo naming
|
||||
# GameCraft weights are already close to FastVideo format with minor adjustments
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# MLP naming: fc1 -> fc_in, fc2 -> fc_out
|
||||
r"^(.*)\.img_mlp\.fc1\.(.*)$": r"\1.img_mlp.fc_in.\2",
|
||||
r"^(.*)\.img_mlp\.fc2\.(.*)$": r"\1.img_mlp.fc_out.\2",
|
||||
r"^(.*)\.txt_mlp\.fc1\.(.*)$": r"\1.txt_mlp.fc_in.\2",
|
||||
r"^(.*)\.txt_mlp\.fc2\.(.*)$": r"\1.txt_mlp.fc_out.\2",
|
||||
r"^(.*)\.img_mlp\.fc1\.(.*)$":
|
||||
r"\1.img_mlp.fc_in.\2",
|
||||
r"^(.*)\.img_mlp\.fc2\.(.*)$":
|
||||
r"\1.img_mlp.fc_out.\2",
|
||||
r"^(.*)\.txt_mlp\.fc1\.(.*)$":
|
||||
r"\1.txt_mlp.fc_in.\2",
|
||||
r"^(.*)\.txt_mlp\.fc2\.(.*)$":
|
||||
r"\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# Single block MLP naming
|
||||
r"^single_blocks\.(\d+)\.mlp\.fc1\.(.*)$": r"single_blocks.\1.mlp.fc_in.\2",
|
||||
r"^single_blocks\.(\d+)\.mlp\.fc2\.(.*)$": r"single_blocks.\1.mlp.fc_out.\2",
|
||||
r"^single_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
|
||||
r"single_blocks.\1.mlp.fc_in.\2",
|
||||
r"^single_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
|
||||
r"single_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# Token refiner naming
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.(.*)$": r"txt_in.refiner_blocks.\1.\2",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.\2",
|
||||
|
||||
# Vector in naming
|
||||
r"^vector_in\.in_layer\.(.*)$": r"vector_in.fc_in.\1",
|
||||
r"^vector_in\.out_layer\.(.*)$": r"vector_in.fc_out.\1",
|
||||
r"^vector_in\.in_layer\.(.*)$":
|
||||
r"vector_in.fc_in.\1",
|
||||
r"^vector_in\.out_layer\.(.*)$":
|
||||
r"vector_in.fc_out.\1",
|
||||
|
||||
# Time embedder naming
|
||||
r"^time_in\.mlp\.0\.(.*)$": r"time_in.mlp.fc_in.\1",
|
||||
r"^time_in\.mlp\.2\.(.*)$": r"time_in.mlp.fc_out.\1",
|
||||
r"^time_in\.mlp\.0\.(.*)$":
|
||||
r"time_in.mlp.fc_in.\1",
|
||||
r"^time_in\.mlp\.2\.(.*)$":
|
||||
r"time_in.mlp.fc_out.\1",
|
||||
|
||||
# Guidance embedder naming (if present)
|
||||
r"^guidance_in\.mlp\.0\.(.*)$": r"guidance_in.mlp.fc_in.\1",
|
||||
r"^guidance_in\.mlp\.2\.(.*)$": r"guidance_in.mlp.fc_out.\1",
|
||||
r"^guidance_in\.mlp\.0\.(.*)$":
|
||||
r"guidance_in.mlp.fc_in.\1",
|
||||
r"^guidance_in\.mlp\.2\.(.*)$":
|
||||
r"guidance_in.mlp.fc_out.\1",
|
||||
|
||||
# Final layer adaLN modulation
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
|
||||
# Refiner block MLP naming
|
||||
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc1\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc2\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# Camera net weights are already correctly named
|
||||
})
|
||||
@@ -118,7 +136,8 @@ class HunyuanGameCraftArchConfig(DiTArchConfig):
|
||||
|
||||
# Layers to exclude from LoRA
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in", "camera_net"])
|
||||
default_factory=lambda:
|
||||
["img_in", "txt_in", "time_in", "vector_in", "camera_net"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
@@ -138,6 +157,7 @@ class HunyuanGameCraftArchConfig(DiTArchConfig):
|
||||
class HunyuanGameCraftConfig(DiTConfig):
|
||||
"""Full config for HunyuanGameCraft model."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=HunyuanGameCraftArchConfig)
|
||||
arch_config: DiTArchConfig = field(
|
||||
default_factory=HunyuanGameCraftArchConfig)
|
||||
|
||||
prefix: str = "HunyuanGameCraft"
|
||||
|
||||
@@ -24,9 +24,12 @@ def is_txt_in(n: str, m) -> bool:
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_refiner_block])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[is_double_block, is_single_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
@@ -87,12 +90,18 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
@@ -116,10 +125,14 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
r"single_blocks.\1.q_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"single_blocks.\1.k_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"single_blocks.\1.linear1.\2", 0, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"single_blocks.\1.linear1.\2", 1, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"single_blocks.\1.linear1.\2", 2, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$": (r"single_blocks.\1.linear1.\2", 3, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 0, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 1, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 2, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 3, 4),
|
||||
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
|
||||
r"single_blocks.\1.linear2.\2",
|
||||
@@ -153,7 +166,8 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
pooled_projection_dim: int = 768
|
||||
rope_theta: int = 256
|
||||
qk_norm: str = "rms_norm"
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -18,9 +18,11 @@ def is_txt_in(n: str, m) -> bool:
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideo15ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
@@ -83,12 +85,18 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
@@ -135,7 +143,8 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
|
||||
target_size: int = 640
|
||||
task_type: str = "i2v"
|
||||
use_meanflow: bool = False
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -18,19 +18,27 @@ def is_txt_in(n: str, m) -> bool:
|
||||
|
||||
@dataclass
|
||||
class HYWorldArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# 1. txt_in submodules (text embedder, refiner blocks):
|
||||
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$": r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$": r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^txt_in\.c_embedder\.linear_1\.(.*)$": r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^txt_in\.c_embedder\.linear_2\.(.*)$": r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$": r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$": r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^txt_in\.c_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^txt_in\.c_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_qkv\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
|
||||
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_proj\.(.*)$":
|
||||
@@ -43,42 +51,66 @@ class HYWorldArchConfig(DiTArchConfig):
|
||||
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
|
||||
# 2. time_in mappings:
|
||||
r"^time_in\.mlp\.0\.(.*)$": r"time_in.timestep_embedder.mlp.fc_in.\1",
|
||||
r"^time_in\.mlp\.2\.(.*)$": r"time_in.timestep_embedder.mlp.fc_out.\1",
|
||||
r"^time_in\.mlp\.0\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_in.\1",
|
||||
r"^time_in\.mlp\.2\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_out.\1",
|
||||
|
||||
# 3. action_in mappings:
|
||||
r"^action_in\.mlp\.0\.(.*)$": r"action_in.mlp.fc_in.\1",
|
||||
r"^action_in\.mlp\.2\.(.*)$": r"action_in.mlp.fc_out.\1",
|
||||
r"^action_in\.mlp\.0\.(.*)$":
|
||||
r"action_in.mlp.fc_in.\1",
|
||||
r"^action_in\.mlp\.2\.(.*)$":
|
||||
r"action_in.mlp.fc_out.\1",
|
||||
|
||||
# 4. byt5_in -> txt_in_2 mappings:
|
||||
r"^byt5_in\.layernorm\.(.*)$": r"txt_in_2.norm.\1",
|
||||
r"^byt5_in\.fc1\.(.*)$": r"txt_in_2.linear_1.\1",
|
||||
r"^byt5_in\.fc2\.(.*)$": r"txt_in_2.linear_2.\1",
|
||||
r"^byt5_in\.fc3\.(.*)$": r"txt_in_2.linear_3.\1",
|
||||
r"^byt5_in\.layernorm\.(.*)$":
|
||||
r"txt_in_2.norm.\1",
|
||||
r"^byt5_in\.fc1\.(.*)$":
|
||||
r"txt_in_2.linear_1.\1",
|
||||
r"^byt5_in\.fc2\.(.*)$":
|
||||
r"txt_in_2.linear_2.\1",
|
||||
r"^byt5_in\.fc3\.(.*)$":
|
||||
r"txt_in_2.linear_3.\1",
|
||||
|
||||
# 5. cond_type_embedding -> cond_type_embed:
|
||||
r"^cond_type_embedding\.(.*)$": r"cond_type_embed.\1",
|
||||
r"^cond_type_embedding\.(.*)$":
|
||||
r"cond_type_embed.\1",
|
||||
|
||||
# 6. vision_in -> image_embedder mappings:
|
||||
r"^vision_in\.proj\.0\.(.*)$": r"image_embedder.norm_in.\1",
|
||||
r"^vision_in\.proj\.1\.(.*)$": r"image_embedder.linear_1.\1",
|
||||
r"^vision_in\.proj\.3\.(.*)$": r"image_embedder.linear_2.\1",
|
||||
r"^vision_in\.proj\.4\.(.*)$": r"image_embedder.norm_out.\1",
|
||||
r"^vision_in\.proj\.0\.(.*)$":
|
||||
r"image_embedder.norm_in.\1",
|
||||
r"^vision_in\.proj\.1\.(.*)$":
|
||||
r"image_embedder.linear_1.\1",
|
||||
r"^vision_in\.proj\.3\.(.*)$":
|
||||
r"image_embedder.linear_2.\1",
|
||||
r"^vision_in\.proj\.4\.(.*)$":
|
||||
r"image_embedder.norm_out.\1",
|
||||
|
||||
# 7. double_blocks mapping:
|
||||
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$": r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$": r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$": r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$": r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 8. Final layer mapping:
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
@@ -118,7 +150,8 @@ class HYWorldArchConfig(DiTArchConfig):
|
||||
ideal_resolution: str = "480p"
|
||||
ideal_task: str = "i2v"
|
||||
task_type: str = "i2v"
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
lambda n, m:
|
||||
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
])
|
||||
|
||||
# Native FastVideo implementation uses the same parameter names as diffusers
|
||||
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
|
||||
# FastVideo uses MLP `fc_in/fc_out`.
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^(.*feed_forward)\.in_layer\.(weight|bias)$": r"\1.mlp.fc_in.\2",
|
||||
r"^(.*feed_forward)\.out_layer\.(weight|bias)$": r"\1.mlp.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Diffusers Kandinsky5Transformer3DModel config fields.
|
||||
in_visual_dim: int = 4
|
||||
in_text_dim: int = 3584
|
||||
in_text_dim2: int = 768
|
||||
time_dim: int = 512
|
||||
out_visual_dim: int = 4
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
model_dim: int = 2048
|
||||
ff_dim: int = 5120
|
||||
num_text_blocks: int = 2
|
||||
num_visual_blocks: int = 32
|
||||
axes_dims: tuple[int, int, int] = (16, 24, 24)
|
||||
visual_cond: bool = False
|
||||
attention_type: str = "regular"
|
||||
attention_causal: bool | None = None
|
||||
attention_local: bool | None = None
|
||||
attention_glob: bool | None = None
|
||||
attention_window: int | None = None
|
||||
attention_P: float | None = None
|
||||
attention_wT: int | None = None
|
||||
attention_wW: int | None = None
|
||||
attention_wH: int | None = None
|
||||
attention_add_sta: bool | None = None
|
||||
attention_method: str | None = None
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
head_dim = sum(self.axes_dims)
|
||||
if self.model_dim % head_dim != 0:
|
||||
raise ValueError(f"model_dim ({self.model_dim}) must be divisible by head_dim ({head_dim})")
|
||||
self.hidden_size = self.model_dim
|
||||
self.num_attention_heads = self.model_dim // head_dim
|
||||
self.in_channels = self.in_visual_dim
|
||||
self.out_channels = self.out_visual_dim
|
||||
self.num_channels_latents = self.in_visual_dim
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5VideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=Kandinsky5ArchConfig)
|
||||
prefix: str = "Kandinsky5"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user