Compare commits

...
Author SHA1 Message Date
mignonjia 8fa6ba6178 mc dfsft 2026-04-01 03:53:20 +00:00
H1yori233 2e5fef787b fix tf scheduler 2026-03-17 18:19:27 -07:00
H1yori233 43d87816bd add logger 2026-03-16 17:27:59 -07:00
H1yori233 2ace7dc6f4 update df scheduler 2026-03-16 15:49:09 -07:00
H1yori233 0ca75db738 fix train / val step mismatch 2026-03-15 20:55:53 -07:00
H1yori233 375ffd3fd5 make visualization in 1 panel 2026-03-15 16:59:34 -07:00
H1yori233 2615ba4291 upload more validation to wandb 2026-03-15 16:43:59 -07:00
RandNMR73 474dd71f28 config 2026-03-15 22:10:50 +00:00
RandNMR73 98ad2d2db6 wangame 2026-03-15 22:01:33 +00:00
alexzms d92858659d [feat] Self-Forcing methods in refactored training infra (#1164) 2026-03-09 20:20:59 -07:00
1a383f3f66 [refactor] train v1 clean up
Co-authored-by: alexzms <3036648523@qq.com>
Co-authored-by: Peiyuan Zhang <a1286225768@slurm-h200-204-215.slurm-compute.tenant-slurm.svc.cluster.local>
2026-03-09 18:58:56 -07:00
alexzmsandPeiyuan Zhang bc27a032c5 [feat] Refactor training framework into fastvideo/train (#1159)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2026-03-09 15:16:42 -07:00
alexzms 2b13e117f0 [Feat] Add causal Wan pipeline with multi-step denoising (#1161) 2026-03-08 13:32:04 -07:00
Junda Chen 99c166c381 feat: Building agent friendly repo (#1151) 2026-03-07 17:46:29 -08:00
XOR-op 95066245db [misc] FlashAttention 4 support (#1114) 2026-03-07 16:53:43 -08:00
Jinzhe Panandgemini-code-assist[bot] 6dcaac768b [CI] PR template (#1157)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-03-07 11:51:28 -08:00
Ajay Anubolu 02c1c49b75 [CI] Add inference performance regression tests (#1140) 2026-03-07 08:26:54 +08:00
Zhang Peiyuan cd1b7cf139 [Refactor] SP Mask --> original seq len; HunyuanVideo 1.5 does not need mask (#1142) 2026-03-04 11:44:47 +08:00
Ajay Anubolu e63b7d8ac4 [Feat] Added OpenAI-compatible API server and benchmark script (#1109) 2026-03-02 17:12:32 -05:00
Jinzhe Pan 5190c1bb1e [Doc] add doc for inference architecture (#1147) 2026-03-02 13:44:27 -08:00
Darren 2cb3bba658 [bugfix]: fix a bug where collect_env was not running properly... (#1145) 2026-03-02 10:57:29 -08:00
Jinzhe Pan f9e1c46c3c [CI][Feat] launch 2 instance to run ssim (#1137) 2026-03-01 01:49:29 -08:00
Peiyuan Zhang e1eda47589 remove temporal frame adjustment 2026-02-27 20:47:16 +00:00
Zhang Peiyuan d902967208 Py/fix sp (#1138) 2026-02-27 12:14:44 +08:00
Zhang Peiyuan fea556269b [Misc] Fix memory leakage in VideoGenerator (#1132) 2026-02-26 19:51:32 -08:00
William Lin 69dd3c68f6 [bugfix] fix matrix game kv indexing and CI (#1135) 2026-02-26 01:24:42 -08:00
Jinzhe Pan 5433f6e80b [fix] preprocessing issue (#1134) 2026-02-25 21:49:52 -08:00
Junda (David) Su e315657066 [docs] [kernel] Migrate to uv (#1127) 2026-02-25 14:13:50 -08:00
Zhang Peiyuan f8d9a0c57f [misc] fix hunyuan (#1125) 2026-02-25 08:26:29 +08:00
Jinzhe PanandWill Lin fa6d276925 [Feat] Improved CI (#1119)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-02-24 12:53:02 -08:00
Zhang PeiyuanandWill Lin fc80d95d7e [Misc] Remove STA (#1124)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-02-23 15:14:42 -08:00
Shao DuanandSolitaryThinker 37cab18780 [bugfix] Added ltx2 guidance missing modulation term (#1100)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-02-23 14:22:31 -08:00
Zhang Peiyuan 128d0b7fc5 [Misc] Remove Teacache (#1121) 2026-02-22 16:53:07 -08:00
Matthew Noto 8092f02e6d small refactor in post-processing to improve efficiency (#1123) 2026-02-22 16:45:44 -08:00
269 changed files with 25716 additions and 763933 deletions
+94
View File
@@ -0,0 +1,94 @@
# Agent Infrastructure — Status Dashboard
Developer-maintained overview of all agent components and their maturity.
Use this to understand what exists, how complete it is, and how much to trust it.
_Last synced: 2026-03-02_
> To resync this dashboard, use the workflow: `.agents/workflows/sync-dashboard.md`
---
## Summary
| Category | Total | ✅ Ready | 🟡 Draft | 🔴 Stub | Trust |
|----------|-------|---------|---------|---------|-------|
| Skills | 8 | 0 | 8 | 0 | Low — newly created, untested |
| Workflows (SOPs) | 4 | 0 | 4 | 0 | Low — newly created, untested |
| Memory files | 4 | 1 | 3 | 0 | Medium — codebase_map is solid |
| Lessons | 0 | — | — | — | N/A — empty |
| Exploration logs | 0 | — | — | — | N/A — empty |
---
## Skills (`.agents/skills/`)
| Skill | File | Status | Trust | Tested | Notes |
|-------|------|--------|-------|--------|-------|
| Launch Experiment | `launch-experiment.md` | 🟡 Draft | Low | ❌ | Needs dry-run validation |
| Monitor Experiment | `monitor-experiment.md` | 🟡 Draft | Low | ❌ | Requires W&B API access to test |
| Summarize Run | `summarize-run.md` | 🟡 Draft | Low | ❌ | Pattern from existing test infra |
| Log Experiment | `log-experiment.md` | 🟡 Draft | Low | ❌ | Journal formatting only |
| Evaluate Video Quality | `evaluate-video-quality.md` | 🟡 Draft | Low | ❌ | SSIM section most mature |
| Index Related Work | `index-related-work.md` | 🟡 Draft | Low | ❌ | Schema defined, no entries yet |
| Search Related Work | `search-related-work.md` | 🟡 Draft | Low | ❌ | Depends on indexed entries |
| Skill Template | `SKILL_TEMPLATE.md` | ✅ Ready | High | ✅ | Meta-template, stable |
### Trust Level Definitions
- **High**: Tested in production, validated against real experiments
- **Medium**: Logic is sound, partially tested or based on existing patterns
- **Low**: Newly written, not yet validated
- **None**: Placeholder only
---
## Workflows / SOPs (`.agents/workflows/`)
| Workflow | File | Status | Trust | Tested | Notes |
|----------|------|--------|-------|--------|-------|
| Experiment Lifecycle | `experiment-lifecycle.md` | 🟡 Draft | Low | ❌ | End-to-end flow, untested |
| Evaluation Development | `evaluation-development.md` | 🟡 Draft | Low | ❌ | Metric dev process |
| Experiment Journaling | `experiment-journaling.md` | 🟡 Draft | Low | ❌ | Journaling cadence |
| Lesson Capture | `lesson-capture.md` | 🟡 Draft | Low | ❌ | Post-experiment reflection |
| Sync Dashboard | `sync-dashboard.md` | 🟡 Draft | Low | ❌ | This dashboard's updater |
---
## Memory (`.agents/memory/`)
| File | Status | Trust | Notes |
|------|--------|-------|-------|
| `codebase_map.md` | ✅ Ready | High | Synthesized from full repo research |
| `experiment_journal.md` | 🟡 Draft | Medium | Schema defined, no entries yet |
| `evaluation_registry.md` | 🟡 Draft | Medium | SSIM/loss metrics documented |
| `related_work/README.md` | 🟡 Draft | Medium | Schema defined, no entries yet |
---
## Lessons (`.agents/lessons/`)
| File | Category | Severity | Notes |
|------|----------|----------|-------|
_No lessons captured yet._
---
## Exploration Logs (`.agents/exploration/`)
| File | Status | Topic | Notes |
|------|--------|-------|-------|
_No exploration logs yet._
---
## What to Do Next
1. **Validate skills**: Run a minimal training experiment using the
`experiment-lifecycle` SOP to test `launch-experiment` → `monitor-experiment`
→ `summarize-run` end-to-end.
2. **Index first related work**: Use `index-related-work` to add at least one
paper (e.g., the Self-Forcing paper used in the codebase).
3. **Capture first lesson**: After the validation run, capture any findings.
4. **Promote to Ready**: As each skill/SOP is tested, update its status here.
+46
View File
@@ -0,0 +1,46 @@
# Exploration Logs
This directory holds draft procedures and investigation notes for tasks that
don't yet have a standardized skill or SOP. Each exploration should follow this
template.
## When to Create an Exploration Log
- You are working on a task with no existing skill or workflow.
- You are experimenting with a new metric, training technique, or tool.
- You want to document findings before they are promoted to a standard.
## File Naming
`<topic-slug>.md` — e.g., `fvd-metric-investigation.md`
## Template
```markdown
# Exploration Log: <Topic>
## Status: draft | under_review | promoted | abandoned
## Context
<Why this exploration is needed — link to experiment or task if applicable.>
## Progress
- [ ] Step 1: ...
- [ ] Step 2: ...
## Findings
<What you have learned so far.>
## Mistakes / Dead Ends
<What didn't work and why — these become lessons.>
## Proposed Standardization
<If this works, describe the skill/SOP/workflow to create.>
```
## Lifecycle
1. **Create** during exploration mode.
2. **Update** as you make progress.
3. **Promote**: If findings are solid, create a skill in `.agents/skills/` or an SOP in `.agents/workflows/`.
4. **Archive mistakes**: Move failures into `.agents/lessons/`.
+48
View File
@@ -0,0 +1,48 @@
# Lessons Learned Database
This directory stores documented mistakes, unexpected behaviors, and their fixes.
Each lesson is a permanent record that helps agents and humans avoid repeating
past errors.
## When to Create a Lesson
- An experiment failed for a non-obvious reason.
- A configuration or hyperparameter choice led to wasted compute.
- A porting, data, or infrastructure issue was discovered and resolved.
- A workaround was needed for a known framework/library bug.
## File Naming
`<YYYY-MM-DD>_<short-slug>.md` — e.g., `2026-03-02_lr-too-high-for-lora.md`
## Template
```markdown
---
date: <ISO-8601>
experiment: <reference to experiment_journal.md entry, if applicable>
category: hyperparameter | data | infrastructure | evaluation | porting | other
severity: critical | important | minor
---
# <Short Descriptive Title>
## What Happened
<Description of the problem and its symptoms.>
## Root Cause
<Analysis of why it happened.>
## Fix / Workaround
<What resolved the issue.>
## Prevention
<How to avoid this in the future — updated skills, SOPs, or checks.>
```
## Usage
- Before starting a task, **search this directory** for relevant lessons.
- After completing or failing a task, **check if a new lesson should be created**.
- Periodically review lessons for **patterns** — recurring themes may warrant
a new skill, SOP, or codebase fix.
+129
View File
@@ -0,0 +1,129 @@
# FastVideo-WorldModel — Codebase Map
High-level structural index for agent orientation. Updated 2026-03-08.
## Repository Layout
```
FastVideo-WorldModel/
├── fastvideo/ # Core Python package
│ ├── models/ # Model implementations
│ │ ├── dits/ # DiT transformers (wanvideo, ltx2, ...)
│ │ ├── vaes/ # VAE models
│ │ ├── encoders/ # Text/image encoders (T5, CLIP)
│ │ ├── schedulers/ # Noise schedulers
│ │ ├── upsamplers/ # Super-resolution models
│ │ ├── audio/ # Audio models
│ │ └── loader/ # Component loaders for HF repos
│ ├── configs/ # Configuration system
│ │ ├── models/ # Arch configs + param_names_mapping
│ │ ├── pipelines/ # Pipeline wiring
│ │ └── sample/ # Default sampling parameters
│ ├── pipelines/ # End-to-end pipelines
│ │ ├── basic/ # Per-model pipelines (wan/, ltx2/, ...)
│ │ └── stages/ # Reusable pipeline stages
│ ├── train/ # Refactored training framework (YAML-driven, preferred)
│ │ ├── trainer.py # Main training loop coordinator
│ │ ├── entrypoint/ # Training entrypoint (train.py) + checkpoint conversion
│ │ ├── methods/ # Training algorithms (FineTune, DFSFT, DMD2, SelfForcing)
│ │ │ ├── base.py # TrainingMethod ABC
│ │ │ ├── fine_tuning/ # FineTuneMethod, DiffusionForcingSFTMethod
│ │ │ └── distribution_matching/ # DMD2Method, SelfForcingMethod
│ │ ├── models/ # Per-role model wrappers (ModelBase, CausalModelBase)
│ │ │ └── wan/ # WanModel, WanCausalModel
│ │ ├── callbacks/ # Composable hooks (grad_clip, ema, validation)
│ │ └── utils/ # Config, builder, checkpoint, optimizer, tracking
│ ├── training/ # Legacy training infrastructure (being phased out)
│ │ ├── trackers.py # W&B tracker (BaseTracker → WandbTracker)
│ │ ├── training_utils.py # Checkpointing, grad clipping, state dicts
│ │ ├── training_pipeline.py # Base training pipeline
│ │ ├── wan_training_pipeline.py # Wan T2V training
│ │ ├── wan_i2v_training_pipeline.py # Wan I2V training
│ │ ├── distillation_pipeline.py # Distillation base
│ │ ├── wan_distillation_pipeline.py # Wan distillation
│ │ ├── self_forcing_distillation_pipeline.py # Self-forcing distill
│ │ ├── ltx2_training_pipeline.py # LTX-2 training
│ │ └── matrixgame_training_pipeline.py # MatrixGame training
│ ├── attention/ # Attention backends
│ ├── distributed/ # Sequence/tensor parallel utilities
│ ├── layers/ # Tensor-parallel layers
│ ├── tests/ # Package-level tests
│ │ ├── training/ # Training regression tests (W&B summary comparison)
│ │ ├── ssim/ # SSIM visual regression tests
│ │ ├── encoders/ # Encoder parity tests
│ │ └── modal/ # Modal CI test runner
│ └── registry.py # Unified config registry
├── fastvideo-kernel/ # CUDA/custom kernels (separate build: ./build.sh)
├── scripts/ # Utility scripts
│ ├── distill/ # Distillation launch scripts
│ ├── inference/ # Inference scripts
│ ├── checkpoint_conversion/ # Weight conversion tools
│ ├── finetune/ # Finetune scripts
│ └── preprocess/ # Data preprocessing
├── examples/ # Ready-to-run examples
│ ├── training/ # Training examples (finetune/, consistency_finetune/)
│ ├── distill/ # Distillation examples
│ ├── inference/ # Inference examples
│ └── dataset/ # Dataset examples
├── docs/ # MkDocs documentation source
│ ├── design/overview.md # Architecture overview
│ ├── training/ # Training guides
│ └── contributing/ # Contributor guides + coding_agents.md
├── tests/ # Top-level tests (local_tests/)
├── AGENTS.md # Agent coding guidelines
└── .agents/ # Agent infrastructure (you are here)
```
## Key Training Entrypoints
### New framework (`fastvideo/train/`) — preferred
| Method | Config Example | Launch Pattern |
|--------|---------------|----------------|
| FineTune (Wan) | `examples/train/finetune_wan2.1_t2v_1.3B_vsa_*.yaml` | `torchrun -m fastvideo.train.entrypoint.train --config <yaml>` |
| DFSFT (Wan causal) | `examples/train/dfsft_wan_causal_t2v_1.3B.yaml` | `torchrun -m fastvideo.train.entrypoint.train --config <yaml>` |
| DMD2 distillation | `examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml` | `torchrun -m fastvideo.train.entrypoint.train --config <yaml>` |
| Self-Forcing | `examples/train/self_forcing_wan_causal_t2v_1.3B.yaml` | `torchrun -m fastvideo.train.entrypoint.train --config <yaml>` |
### Legacy pipelines (`fastvideo/training/`) — being phased out
| Pipeline | Entrypoint | Launch Pattern |
|----------|-----------|----------------|
| Wan T2V finetune | `fastvideo/training/wan_training_pipeline.py` | `torchrun --nproc_per_node N` |
| Wan I2V finetune | `fastvideo/training/wan_i2v_training_pipeline.py` | `torchrun --nproc_per_node N` |
| Wan distillation (DMD) | `fastvideo/training/wan_distillation_pipeline.py` | `torchrun --nproc_per_node N` |
| Self-forcing distill | `fastvideo/training/wan_self_forcing_distillation_pipeline.py` | `torchrun --nproc_per_node N` |
| LTX-2 finetune | `fastvideo/training/ltx2_training_pipeline.py` | `torchrun --nproc_per_node N` |
| MatrixGame | `fastvideo/training/matrixgame_training_pipeline.py` | `torchrun --nproc_per_node N` |
## W&B Integration
- **Tracker classes**: `fastvideo/training/trackers.py`
- `WandbTracker` — logs metrics, videos, timing
- `SequentialTracker` — fan-out to multiple trackers
- `DummyTracker` — no-op for offline/test
- **Run summary location**: `<output_dir>/tracker/wandb/latest-run/files/wandb-summary.json`
- **Reference summaries**: `fastvideo/tests/training/*/` (e.g., `a40_reference_wandb_summary.json`)
- **Environment**: `WANDB_API_KEY`, `WANDB_BASE_URL`, `WANDB_MODE`
## Critical Environment Variables
| Variable | Purpose |
|----------|---------|
| `WANDB_API_KEY` | W&B authentication |
| `WANDB_MODE` | `online` / `offline` |
| `FASTVIDEO_ATTENTION_BACKEND` | `FLASH_ATTN` / `TORCH_SDPA` |
| `TOKENIZERS_PARALLELISM` | Set `false` to avoid fork warnings |
| `HF_HOME` | HuggingFace cache directory |
## Build & Test Commands
```bash
uv pip install -e .[dev] # Editable install
pre-commit run --all-files # Lint/format/spell
pytest tests/ # Top-level tests
pytest fastvideo/tests/ -v # Package tests
pytest fastvideo/tests/training/Vanilla -srP # Training loss regression
pytest fastvideo/tests/ssim/ -vs # SSIM visual regression
cd fastvideo-kernel && ./build.sh # Build kernels
```
@@ -0,0 +1,327 @@
# Evaluation Metrics Registry
Living catalog of all evaluation metrics for FastVideo-WorldModel video quality
assessment. Each metric includes a detailed explanation, implementation status,
usage instructions, and interpretation guide.
_Last updated: 2026-03-02_
---
## Metric Summary
| Metric | Category | Status | Location | Trust |
|--------|----------|--------|----------|-------|
| **FVD** | Distribution | ✅ Implemented | `benchmarks/fvd/` | High |
| **SSIM** | Reference | ✅ Implemented | `fastvideo/tests/ssim/` | High |
| **LPIPS** | Perceptual | ✅ Implemented | `scripts/lora_extraction/` | Medium |
| **Loss trajectory** | Training signal | ✅ Implemented | W&B `train_loss` | Medium |
| **Grad norm stability** | Training signal | ✅ Implemented | W&B `grad_norm` | Medium |
| **GameWorld Score** | Multi-dim benchmark | 🟡 External | Matrix-Game repo | Low |
| **Human preference** | Gold standard | 🔴 Manual | N/A | Highest |
---
## Implemented Metrics
### FVD — Fréchet Video Distance
**Category**: Distribution-level quality metric
**Status**: ✅ Fully implemented in `benchmarks/fvd/`
**Trust**: High — standard protocol, I3D feature extractor
#### What It Measures
FVD measures the distance between the **distribution** of generated videos and
a distribution of real/reference videos. It works by:
1. Extracting spatiotemporal features from both real and generated video sets
using a pretrained **I3D** (Inflated 3D ConvNet) model.
2. Modeling each set of features as a multivariate Gaussian (mean + covariance).
3. Computing the **Fréchet distance** between the two Gaussians.
Lower FVD = generated videos are more statistically similar to real videos.
#### Why It Matters
- FVD is the **de facto standard** for benchmarking video generation models.
- It captures both **visual quality** (are individual frames realistic?) and
**temporal coherence** (do frames flow naturally?).
- Matrix-Game 2.0, Open-Sora, and most video generation papers report FVD.
#### Limitations
- Requires a **large sample set** (standard protocol uses 2048 videos) to
produce stable statistics. Small sample sizes yield noisy results.
- Measures **distributional similarity**, not per-video quality. A model could
have low FVD by generating a diverse set of "roughly okay" videos.
- The I3D model was trained on Kinetics-400 (human actions). It may be less
sensitive to domain-specific artifacts in non-human-action videos (e.g.,
driving, game environments).
- Does not directly measure text-video alignment or action controllability.
#### How to Use
```python
# Programmatic
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
config = FVDConfig.fvd2048_16f() # Standard: 2048 videos, 16 frames
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
print(f"FVD: {results['fvd']:.2f}")
```
```bash
# CLI
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f
```
**Preset protocols**:
| Protocol | Videos | Frames | Use Case |
|----------|--------|--------|----------|
| `fvd2048_16f` | 2048 | 16 | Standard benchmark (papers) |
| `fvd2048_128f` | 2048 | 128 | Long video evaluation |
| `quick_test` | 100 | 16 | Fast dev iteration |
**Feature extractors**: `i3d` (default, standard), `clip`, `videomae`
#### Interpretation
| FVD Range | Interpretation |
|-----------|---------------|
| < 100 | Excellent — near-real quality |
| 100–300 | Good — competitive with SOTA |
| 300–600 | Fair — noticeable gap from real |
| > 600 | Poor — significant quality issues |
> FVD values are dataset-dependent. Always compare against baselines evaluated
> on the same real video distribution.
---
### SSIM — Structural Similarity Index
**Category**: Per-frame reference comparison
**Status**: ✅ Implemented in `fastvideo/tests/ssim/`
**Trust**: High — used in CI regression tests
#### What It Measures
SSIM compares two images (or video frames) based on three components:
1. **Luminance**: brightness similarity
2. **Contrast**: dynamic range similarity
3. **Structure**: spatial pattern similarity
The final score is a value in [0, 1] where 1.0 = identical.
#### Why It Matters
- Used as a **regression guard** in CI: ensures model updates don't degrade
visual output below a threshold.
- More perceptually meaningful than raw pixel MSE.
- Fast to compute — suitable for automated testing.
#### Limitations
- Requires a **pixel-aligned reference** video. Cannot compare videos with
different seeds, prompts, or angles.
- Operates **per-frame** — does not capture temporal coherence.
- Insensitive to some perceptual artifacts (color shifts, high-frequency noise).
#### How to Use
```bash
pytest fastvideo/tests/ssim/ -vs
```
#### Interpretation
| SSIM Range | Quality |
|------------|---------|
| > 0.90 | Excellent — very close to reference |
| 0.80–0.90 | Good — acceptable for most uses |
| 0.70–0.80 | Fair — noticeable differences |
| < 0.70 | Poor — significant divergence |
---
### LPIPS — Learned Perceptual Image Patch Similarity
**Category**: Per-frame perceptual distance
**Status**: ✅ Implemented in `scripts/lora_extraction/lora_inference_comparison.py`
**Trust**: Medium — available but only used for LoRA comparison currently
#### What It Measures
LPIPS uses a pretrained neural network (AlexNet by default) to extract
deep features from two images and computes the distance between them in
feature space. Unlike SSIM, LPIPS correlates much more strongly with
**human perceptual judgments**.
Lower LPIPS = more perceptually similar.
#### Why It Matters
- Best available automated proxy for **human visual judgments** at the frame
level.
- Captures semantic and structural differences that SSIM misses (e.g., texture
changes, minor recoloring).
- Used for validating LoRA merge quality.
#### Limitations
- Per-frame metric — no temporal awareness.
- Requires reference video (paired comparison only).
- Slightly slower than SSIM due to neural network forward pass.
#### How to Use
```bash
python scripts/lora_extraction/lora_inference_comparison.py \
--base merged_model \
--ft path/to/finetuned \
--adapter NONE \
--output-dir results \
--prompt "A cat" \
--compute-lpips
```
#### Interpretation
| LPIPS Range | Quality |
|-------------|---------|
| < 0.10 | Excellent — nearly indistinguishable |
| 0.10–0.20 | Good — minor perceptual differences |
| 0.20–0.40 | Fair — noticeable differences |
| > 0.40 | Poor — clearly different |
---
### Loss Trajectory
**Category**: Training signal proxy
**Status**: ✅ Active (from W&B `train_loss`)
**Trust**: Medium — proxy, not direct quality measure
#### What It Measures
Tracks the training loss over time. A healthy training run shows:
- **Decreasing loss** over the first hundreds of steps.
- **Stable gradient norms** (no wild spikes).
- **Consistent step times** (no infrastructure issues).
#### Why It Matters
- Cheapest evaluation signal — available in real-time from W&B.
- Critical for the **30-minute quality check** workflow.
- At later training stages (when loss becomes meaningful), trajectory shape
can predict final model quality.
#### Context: How This Evolves
The team's experience shows evaluation signals change during a project:
- **Early stage**: Loss may be flat or meaningless → focus on SSIM & visual
inspection instead.
- **Mid stage**: Loss starts decreasing → trajectory shape becomes useful.
- **Late stage**: Loss is meaningful → can compare trajectories across runs.
This dynamic is a key insight from the team's workflow: don't over-rely on
loss early; don't ignore it late.
---
### Grad Norm Stability
**Category**: Training health diagnostic
**Status**: ✅ Active (from W&B `grad_norm`)
**Trust**: Medium — diagnostic, not quality metric
#### What It Measures
The magnitude of gradients during training. Stable grad norms indicate
healthy optimization. Spikes or NaN values indicate training instability.
#### Alert Thresholds
| Condition | Meaning |
|-----------|---------|
| Stable ~0.3–0.5 | Normal training |
| Single spike > 3× average | Possible bad batch, monitor |
| NaN or Inf | 🔴 Training has diverged — stop run |
| Increasing trend | Learning rate may be too high |
---
## External Benchmarks
### GameWorld Score Benchmark (Matrix-Game)
**Category**: Multi-dimensional evaluation framework for interactive world models
**Status**: 🟡 External — not implemented in-repo
**Source**: [Matrix-Game 1.0 benchmark](https://github.com/SkyworkAI/Matrix-Game), used in [Matrix-Game 2.0 paper](https://arxiv.org/abs/2508.13009)
#### What It Measures
A comprehensive benchmark examining **four critical capabilities**:
| Dimension | What It Evaluates | Example Signals |
|-----------|-------------------|-----------------|
| **Visual quality** | Frame-level realism, absence of artifacts | Color fidelity, sharpness, coherence |
| **Temporal quality** | Smoothness across frames, motion consistency | Jitter, flickering, temporal aliasing |
| **Action controllability** | Response to input actions (keyboard/mouse) | Action delay, correctness, smoothness |
| **Physical rule understanding** | Adherence to physics (gravity, collision) | Object persistence, plausible motion |
#### Context from Matrix-Game 2.0
- Evaluation uses **597-frame composite action sequences** over 32 Minecraft
scenes and 16 wild scenes.
- Action controllability assessment is **Minecraft-specific** — cannot be
directly applied to wild/general scenes.
- The paper notes that models that "collapse" to static frames can
paradoxically score higher on consistency metrics — beware of this confound.
#### Relevance to FastVideo
- Matrix-Game 2.0 is built on SkyReels-V2/Wan2.1 architecture — **same model
family as FastVideo**.
- Their distillation uses DMD-based Self-Forcing — **same technique** as our
`self_forcing_distillation_pipeline.py`.
- GameWorld Score dimensions are a useful framework for thinking about world
model quality even outside gaming contexts.
---
## Human Preference Evaluation
**Category**: Gold-standard quality assessment
**Status**: 🔴 Manual process — no automated implementation
**Priority**: **Highest** — this is the most important evaluation signal
**Trust**: Highest — but expensive
### What It Measures
Human evaluators compare generated videos and rate them on dimensions like:
- Overall quality and realism
- Temporal coherence and smoothness
- Prompt adherence / action correctness
- Absence of artifacts
#### Why It's the Most Important Metric
All automated metrics are **proxies** for human judgment. They can be gamed
or may miss artifacts that humans easily notice. Human preference is the
ultimate ground truth for video generation quality.
#### Cost & Practicality
| Approach | Cost | Scale | When to Use |
|----------|------|-------|-------------|
| Internal team review | Low | ~10–50 videos | Every major checkpoint |
| Crowdsource (MTurk, Scale) | Medium | 100+ videos | Pre-release validation |
| A/B preference test | Medium | Pairs | Comparing two model versions |
#### Recommended Protocol
1. Sample 10–20 videos from the model at a checkpoint.
2. Include diverse prompts (easy + hard, short + long).
3. Have 2–3 evaluators score each video 1–5 on: quality, coherence, fidelity.
4. Record scores in the experiment journal.
---
## Metrics NOT Used
| Metric | Reason |
|--------|--------|
| ~~CLIP-Score~~ | Not used by the team. Measures text-image alignment using CLIP embeddings, but not well-suited for video temporal quality. |
| Inception Score (IS) | Less informative than FVD for video; primarily an image metric. |
| PSNR | Pixel-level metric; less perceptually meaningful than SSIM/LPIPS. |
---
## Adding a New Metric
Follow the SOP: `.agents/workflows/evaluation-development.md`
1. Prototype in `.agents/exploration/`
2. Validate on known-good and known-bad samples
3. Add to this registry
4. Update the `evaluate-video-quality` skill
@@ -0,0 +1,21 @@
# Experiment Journal
Living log of all experiments. Each entry captures what was tried, the result,
and any insights. Newest entries go at the top.
_No experiments logged yet. Use the `log-experiment` skill to add entries._
<!-- TEMPLATE — copy and fill for each new experiment:
## [YYYY-MM-DD] Experiment: <name>
- **Hypothesis**: <what you expected to learn>
- **Config**: model=..., lr=..., sp_size=..., gpus=..., script=...
- **W&B run**: <run_id or URL>
- **Duration**: <total wall time>
- **Key metrics**: loss=..., step_time=..., grad_norm=...
- **Checkpoint**: <path>
- **Insight**: <what was learned>
- **Status**: running | completed | failed | abandoned
- **Related lessons**: `.agents/lessons/<filename>.md`
-->
+4
View File
@@ -0,0 +1,4 @@
{"name": "codebase-map", "description": "High-level structural index of the FastVideo-WorldModel repository", "path": "codebase-map/README.md", "status": "ready", "trust": "high"}
{"name": "evaluation-registry", "description": "Catalog of all evaluation metrics with detailed explanations, implementation status, and usage guides", "path": "evaluation-registry/README.md", "status": "draft", "trust": "medium"}
{"name": "experiment-journal", "description": "Living log of all experiments with hypotheses, configs, metrics, and insights", "path": "experiment-journal/README.md", "status": "draft", "trust": "medium"}
{"name": "related-work", "description": "Index of related papers, repos, and blog posts with structured comparisons to FastVideo", "path": "related-work/README.md", "status": "draft", "trust": "low"}
+34
View File
@@ -0,0 +1,34 @@
# Related Work Index
Each file in this directory is a structured summary of a related paper, repo,
or blog post relevant to FastVideo-WorldModel training.
## File Format
Each file is named `<slug>.md` and follows this structure:
```markdown
---
title: <paper/repo title>
source: <URL or citation>
type: paper | repo | blog
date_indexed: <ISO-8601>
tags: [world-model, distillation, evaluation, reward-shaping, ...]
---
## Summary
<1-2 paragraph summary of the work.>
## Key Differences from FastVideo
- <Bullet points comparing their approach to ours.>
## Actionable Insights
- <What we could adopt or adapt.>
```
## How to Add New Entries
Use the `index-related-work` skill, or manually create a file following the
template above.
_No related work indexed yet._
+76
View File
@@ -0,0 +1,76 @@
# Agent Onboarding — FastVideo-WorldModel
Welcome, agent. This is the **master onboarding** guide. Follow the steps below,
then check if a **domain-specific onboarding** exists for your task.
## Domain-Specific Onboarding
If your task falls into one of these areas, read the specialized guide **after**
completing the general steps below:
| Domain | Guide | When to Use |
|--------|-------|-------------|
| **WorldModel Training** | `worldmodel-training/README.md` | Training, finetuning, distillation, experiment management |
---
## Step 1: Understand the Codebase
Read these files to build your context:
| Priority | File | What you learn |
|----------|------|----------------|
| 1 | `AGENTS.md` | Coding guidelines, build/test commands, PR conventions |
| 2 | `docs/design/overview.md` | Architecture: models, pipelines, configs, registry |
| 3 | `fastvideo/train/` | Refactored training framework (YAML-driven, modular methods/models/callbacks) |
| 4 | `docs/training/overview.md` | Training data flow and preprocessing |
| 5 | `docs/training/finetune.md` | Training arguments, parallelism, LoRA, validation |
| 6 | `docs/contributing/coding_agents.md` | How to add model pipelines with agent assistance |
## Step 2: Discover Available Resources
Read these two index files to see what skills and memory modules exist:
- **`.agents/skills/index.jsonl`** — catalog of all agent skills (name + description)
- **`.agents/memory/index.jsonl`** — catalog of all memory modules (name + description)
Each entry has a `path` field pointing to the full content. Only load the
full README.md for modules relevant to your current task.
## Step 3: Check for Existing Skills & SOPs
Before writing new code or procedures:
1. **Skills**: Read `.agents/skills/index.jsonl` — find a matching skill by description.
2. **Workflows/SOPs**: Browse `.agents/workflows/` — step-by-step procedures for common tasks.
3. **Lessons**: Browse `.agents/lessons/` — known pitfalls and their fixes.
If a skill or SOP exists for your task, **use it**. If not, you are in **exploration mode** — see Step 4.
## Step 4: Exploration Mode
If no existing skill/SOP covers your task:
1. Document your progress in `.agents/exploration/<topic>.md` using the template in `.agents/exploration/README.md`.
2. At the end of your session, reflect:
- **What worked** → propose a new skill or SOP in the exploration log.
- **What failed** → create a lesson in `.agents/lessons/`.
3. Flag the exploration log for human review.
## Quick Reference
```
.agents/
├── ONBOARDING.md ← you are here
├── STATUS.md ← dashboard: completeness & trust of all components
├── skills/ ← reusable agent skills
├── workflows/ ← SOPs and procedures
├── memory/ ← persistent context (folder per topic + index.jsonl)
│ ├── index.jsonl
│ ├── codebase-map/
│ ├── experiment-journal/
│ ├── evaluation-registry/
│ └── related-work/
├── lessons/ ← mistakes and fixes
└── exploration/ ← draft procedures
```
@@ -0,0 +1,302 @@
# WorldModel Training — Agent Onboarding
Specialized onboarding for agents working on FastVideo-WorldModel training,
distillation, and evaluation. Read the master onboarding (`.agents/onboarding/README.md`)
first, then come here.
---
## Domain Context
FastVideo-WorldModel trains **interactive world models** — video generation systems
that respond to user actions (keyboard/mouse) in real-time. The architecture is
based on **Wan2.1** (SkyReels-V2) DiT models with causal attention for
auto-regressive streaming generation.
**Key techniques you will work with:**
- Full finetuning and LoRA on Wan / LTX-2 / MatrixGame models
- DMD-based distillation (few-step generation)
- Self-Forcing distillation (causal streaming)
- Diffusion-Forcing SFT (DFSFT) for causal models
- VSA (Variable Sparsity Acceleration) for efficient training
---
## Training Code: Two Generations
### New modular framework: `fastvideo/train/` (preferred)
The refactored training code uses a **YAML-only config-driven** architecture
with composable methods, per-role models, and a callback system. All new
training work should use this framework.
### Legacy pipelines: `fastvideo/training/` (deprecated)
The old monolithic pipeline classes (`WanTrainingPipeline`,
`DistillationPipeline`, etc.) still exist but are being phased out. The new
framework imports select utilities from `fastvideo/training/` for backward
compatibility (EMA, gradient clipping, checkpoint wrappers).
---
## Essential Reading (Training-Specific)
Read these **in order** before touching any training code:
| # | File | What You Learn |
|---|------|----------------|
| 1 | `docs/training/overview.md` | Training data flow: raw video → text embeddings + video latents → training |
| 2 | `docs/training/finetune.md` | Training arguments, parallelism (SP/TP), LoRA, validation settings |
| 3 | `docs/training/data_preprocess.md` | How to preprocess datasets into the expected format |
| 4 | `docs/design/overview.md` | Architecture: models, pipelines, configs, registry |
---
## New Training Framework (`fastvideo/train/`)
### Architecture Overview
```
fastvideo/train/
├── __init__.py → exports Trainer
├── trainer.py → main training loop coordinator
├── entrypoint/
│ ├── train.py → YAML-only training entrypoint
│ └── dcp_to_diffusers.py → checkpoint conversion utility
├── methods/ → training algorithms (TrainingMethod ABC)
│ ├── base.py → TrainingMethod base class
│ ├── fine_tuning/
│ │ ├── finetune.py → FineTuneMethod (supervised finetuning)
│ │ └── dfsft.py → DiffusionForcingSFTMethod (causal)
│ ├── distribution_matching/
│ │ ├── dmd2.py → DMD2Method (distribution matching distill)
│ │ └── self_forcing.py → SelfForcingMethod (causal streaming)
│ ├── knowledge_distillation/ → (stub, not yet implemented)
│ └── consistency_model/ → (stub, not yet implemented)
├── models/ → per-role model instances
│ ├── base.py → ModelBase & CausalModelBase (ABC)
│ └── wan/
│ ├── wan.py → WanModel (non-causal)
│ └── wan_causal.py → WanCausalModel (causal streaming)
├── callbacks/ → training hooks & monitoring
│ ├── callback.py → Callback base class + CallbackDict
│ ├── grad_clip.py → GradNormClipCallback
│ ├── ema.py → EMACallback (shadow weights)
│ └── validation.py → ValidationCallback (sampling + eval)
└── utils/ → configuration, building, checkpointing
├── builder.py → build_from_config() (config → runtime)
├── checkpoint.py → CheckpointManager (DCP-based)
├── config.py → load_run_config() (YAML → RunConfig)
├── training_config.py → TypedConfig dataclasses
├── optimizer.py → build_optimizer_and_scheduler()
├── instantiate.py → resolve_target() + instantiate()
├── tracking.py → build_tracker() (W&B, etc.)
├── dataloader.py → dataloader utilities
├── module_state.py → apply_trainable()
└── moduleloader.py → load_module_from_path()
```
### Key Concepts
**TrainingMethod** (`methods/base.py`): Abstract base class for all training
algorithms. Owns role models (student, teacher, critic), manages checkpoint
state, and defines the training step interface.
**ModelBase** (`models/base.py`): Per-role model wrapper. Each role (student,
teacher, critic) gets its own `ModelBase` instance owning a `transformer` and
`noise_scheduler`. `CausalModelBase` extends this for streaming models.
**Callback system** (`callbacks/`): Composable hooks for gradient clipping,
EMA, validation, etc. Configured via YAML, dispatched by `CallbackDict`.
**Config system** (`utils/config.py`, `utils/training_config.py`): YAML files
are parsed into typed `RunConfig` dataclass trees. Models and methods use
`_target_` fields for instantiation (similar to Hydra).
### Training Flow
```
run_training_from_config(config_path)
→ load_run_config() # YAML → RunConfig
→ init_distributed() # TP/SP setup
→ build_from_config() # instantiate models, method, dataloader
→ Trainer.run() # main loop:
├─ callbacks.on_train_start()
├─ checkpoint_manager.maybe_resume()
├─ for step in range(max_steps):
│ ├─ method.single_train_step(batch)
│ ├─ method.backward()
│ ├─ callbacks.on_before_optimizer_step()
│ ├─ method.optimizers_schedulers_step()
│ ├─ tracker.log(metrics, step)
│ ├─ callbacks.on_training_step_end()
│ └─ checkpoint_manager.maybe_save(step)
├─ callbacks.on_train_end()
└─ checkpoint_manager.save_final()
```
### Training Methods
| Method | Class | Use Case |
|--------|-------|----------|
| **FineTune** | `FineTuneMethod` | Single-role supervised finetuning |
| **DFSFT** | `DiffusionForcingSFTMethod` | Diffusion-forcing SFT with inhomogeneous timesteps |
| **DMD2** | `DMD2Method` | Multi-role distribution matching distillation (student + teacher + critic) |
| **Self-Forcing** | `SelfForcingMethod` | Extends DMD2 for causal student rollouts |
### Launching Training (New Framework)
Training is launched via `torchrun` with a single YAML config:
```bash
torchrun --nproc_per_node <N_GPUS> \
-m fastvideo.train.entrypoint.train \
--config examples/train/<config>.yaml
```
### Example YAML Configs
| Config | Method | Description |
|--------|--------|-------------|
| `examples/train/finetune_wan2.1_t2v_1.3B_vsa_phase3.4_0.9sparsity.yaml` | FineTune | Wan 1.3B finetuning with VSA sparsity |
| `examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml` | DMD2 | Wan 1.3B distillation (student + teacher + critic) |
| `examples/train/dfsft_wan_causal_t2v_1.3B.yaml` | DFSFT | Causal Wan 1.3B diffusion-forcing SFT |
| `examples/train/self_forcing_wan_causal_t2v_1.3B.yaml` | Self-Forcing | Causal streaming distillation |
### Checkpointing (New Framework)
**CheckpointManager** (`utils/checkpoint.py`) saves via `torch.distributed.checkpoint`:
```
output_dir/
└─ checkpoint-{step}/
├─ dcp/ # DCP state dict
├─ config.json # resolved training config
└─ .fastvideo_metadata.json
```
Checkpoint state includes: role model weights, per-role optimizers/schedulers,
CUDA RNG state, and callback state (e.g., EMA shadow weights).
### Config Structure
A YAML config defines the full training pipeline:
```yaml
models:
student:
_target_: fastvideo.train.models.wan.WanModel
model_path: ...
trainable: true
teacher: # optional, for distillation
_target_: fastvideo.train.models.wan.WanModel
model_path: ...
trainable: false
method:
_target_: fastvideo.train.methods.fine_tuning.FineTuneMethod
# method-specific params...
training:
distributed: { num_gpus: 8, tp_size: 1, sp_size: 8 }
data: { data_path: ..., batch_size: 1 }
optimizer: { lr: 1e-5, lr_scheduler: constant_with_warmup }
loop: { max_train_steps: 1000 }
checkpoint: { output_dir: ./outputs }
tracker: { trackers: [wandb], project_name: ... }
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.ValidationCallback
validation_steps: 100
```
---
## Legacy Training Pipelines (`fastvideo/training/`)
> **Note:** Use the new `fastvideo/train/` framework for new work. This section
> is retained for reference on existing pipelines not yet migrated.
| Pipeline | Entrypoint | Use Case |
|----------|-----------|----------|
| Wan T2V finetune | `fastvideo/training/wan_training_pipeline.py` | Standard text-to-video finetune / LoRA |
| Wan I2V finetune | `fastvideo/training/wan_i2v_training_pipeline.py` | Image-to-video (first frame conditioned) |
| MatrixGame finetune | `fastvideo/training/matrixgame_training_pipeline.py` | Action-conditioned world model |
| LTX-2 finetune | `fastvideo/training/ltx2_training_pipeline.py` | LTX-2 architecture finetuning |
| Wan DMD distillation | `fastvideo/training/wan_distillation_pipeline.py` | Few-step distillation via DMD |
| Self-Forcing distill | `fastvideo/training/wan_self_forcing_distillation_pipeline.py` | Causal streaming distillation |
---
## Key Infrastructure
### W&B Integration
- **Tracker**: `fastvideo/training/trackers.py` — `WandbTracker` class
- **New framework tracker**: `fastvideo/train/utils/tracking.py` — `build_tracker()`
- **Env vars**: `WANDB_API_KEY`, `WANDB_BASE_URL`, `WANDB_MODE`
### Parallelism
- **SP** (Sequence Parallel): splits video frames across GPUs — `sp_size: N`
- **TP** (Tensor Parallel): splits model layers across GPUs — `tp_size: N`
- Typical configs: SP=2–8, TP=1–2
---
## Evaluation (for training runs)
Read `.agents/memory/evaluation-registry/README.md` for the full metric catalog.
**Quick summary for training agents:**
| Metric | When to Use | Trust |
|--------|-------------|-------|
| **Loss trajectory** | Every run, real-time from W&B | Medium |
| **SSIM** | When comparing against reference outputs | High |
| **FVD** | For benchmarking model quality (`benchmarks/fvd/`) | High |
| **LPIPS** | LoRA merge validation | Medium |
| **Human preference** | Major checkpoints | Highest |
---
## Common Workflows
| Task | Skill / SOP |
|------|-------------|
| Launch a training run | `.agents/skills/launch-experiment/SKILL.md` |
| Monitor a running experiment | `.agents/skills/monitor-experiment/SKILL.md` |
| Summarize final results | `.agents/skills/summarize-run/SKILL.md` |
| Full experiment lifecycle | `.agents/workflows/experiment-lifecycle.md` |
| Capture lessons from failures | `.agents/workflows/lesson-capture.md` |
---
## World Model–Specific Concepts
### Action Injection (MatrixGame)
The MatrixGame pipeline adds **action modules** to each DiT block, enabling
frame-level mouse/keyboard input conditioning. The action sequence is injected
per-frame alongside the latent video tokens.
### Causal Architecture
For streaming generation, the model uses **causal attention** (each frame only
attends to previous frames). This enables auto-regressive chunk-by-chunk
generation — critical for real-time interactive world models.
### Self-Forcing Distillation
A **data-free** distillation method where the student model is trained to
generate coherent video sequences by being forced to use its own previous
outputs (rather than ground-truth) as context. This produces models robust to
their own error accumulation during long auto-regressive generation.
### DMD Distillation (Distribution Matching Distillation)
Reduces inference steps from ~50 to 3–4 by training a student model to match
the output distribution of the teacher model. Uses a critic network to estimate
distribution divergence.
### Diffusion-Forcing SFT (DFSFT)
Supervised finetuning with **inhomogeneous timesteps** across chunks — each
chunk in a causal sequence can have a different noise level, training the model
to handle mixed-fidelity contexts.
+57
View File
@@ -0,0 +1,57 @@
---
name: <skill-name>
description: <one-line description — Codex uses this for implicit invocation matching>
---
# <Skill Name>
## Purpose
<Why this skill exists and when to use it.>
## Prerequisites
- <What must be true before using this skill>
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `param1` | Yes | ... |
## Steps
1. **Step 1 title**
- Detail...
2. **Step 2 title**
- Detail...
## Outputs
- <What this skill produces>
## Example Usage
```
<Example invocation or prompt snippet>
```
## References
- <Links to relevant files in the codebase>
---
## Folder Structure
Each skill lives in its own directory under `.agents/skills/`:
```
.agents/skills/<skill-name>/
├── SKILL.md # Required: instructions + metadata (this file)
├── scripts/ # Optional: executable helper scripts
├── references/ # Optional: documentation, papers
└── assets/ # Optional: templates, resources
```
After creating a new skill, add an entry to `.agents/skills/index.jsonl`:
```json
{"name": "<skill-name>", "description": "<description>", "path": "<skill-name>/SKILL.md", "status": "draft", "trust": "low"}
```
@@ -0,0 +1,128 @@
---
name: evaluate-video-quality
description: Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)
---
# Evaluate Video Quality
## Purpose
Assess the quality of videos generated by a training run. Combines multiple
signals to give a holistic quality assessment. This skill is **evolving** —
new metrics will be added as they are developed.
## Prerequisites
- Generated videos available locally or via W&B artifacts.
- For SSIM: reference videos from official implementations.
- For caption consistency: LLM access (optional, stub for now).
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `video_paths` | Yes | List of paths to generated videos |
| `reference_paths` | No | Paths to reference videos (for SSIM) |
| `prompts` | No | Prompts used to generate videos (for caption check) |
| `loss_summary` | No | Path to W&B summary JSON (for loss trajectory) |
| `metrics` | No | Which metrics to run (default: all available) |
## Available Metrics
Check `.agents/memory/evaluation-registry/README.md` for the current catalog.
### SSIM (Active)
Leverages the existing infrastructure in `fastvideo/tests/ssim/`.
```bash
pytest fastvideo/tests/ssim/ -vs --video-path <generated> --reference-path <reference>
```
Or use the SSIM utility directly:
```python
from fastvideo.tests.ssim.ssim_utils import compute_ssim
score = compute_ssim(generated_video, reference_video)
# score > 0.85 is typically "acceptable"
```
**Interpretation**:
| SSIM Range | Quality |
|------------|---------|
| > 0.90 | Excellent — very close to reference |
| 0.80–0.90 | Good — acceptable for most uses |
| 0.70–0.80 | Fair — noticeable differences |
| < 0.70 | Poor — significant quality issues |
### Loss Trajectory (Active)
Analyze the loss curve shape from W&B summary:
```python
import json
with open(loss_summary_path) as f:
summary = json.load(f)
final_loss = summary["train_loss"]
runtime = summary["_runtime"]
steps = summary["_step"]
```
**Early-stage heuristics** (first 500 steps):
- Loss should be decreasing (even slightly).
- Grad norm should be stable (no wild oscillations).
- If loss is flat or increasing, flag for review.
### Caption Consistency (Draft — Not Yet Calibrated)
Use an LLM to evaluate whether the video content matches the input prompt.
```
Prompt: "A golden retriever playing in the snow"
Video: <path>
Score the video on:
1. Object presence (is there a golden retriever?)
2. Action accuracy (is it playing?)
3. Environment match (is there snow?)
4. Overall coherence (does it look natural?)
Each 1-5, total /20.
```
> ⚠️ This metric is in **draft** status. Results should not be treated as
> ground truth until calibrated against human judgments.
## Steps
1. **Identify available metrics** — Check `.agents/memory/evaluation-registry/README.md`.
2. **Run each metric** — Collect scores.
3. **Aggregate** — Produce a combined quality report.
4. **Log** — Update the experiment journal with quality results.
## Outputs
```markdown
## Video Quality Report: <experiment_name>
| Metric | Score | Threshold | Status |
|--------|-------|-----------|--------|
| SSIM (avg) | 0.87 | > 0.80 | ✅ Pass |
| Loss trajectory | decreasing | decreasing | ✅ Pass |
| Caption consistency | 16/20 | > 14/20 | ✅ Pass |
### Per-Video Scores
| Video | SSIM | Caption |
|-------|------|---------|
| video_001.mp4 | 0.89 | 17/20 |
| video_002.mp4 | 0.85 | 15/20 |
```
## References
- `fastvideo/tests/ssim/` — SSIM test infrastructure
- `fastvideo/tests/training/Vanilla/test_training_loss.py` — loss comparison
- `.agents/memory/evaluation-registry/README.md` — metric catalog
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version with SSIM, loss trajectory, caption consistency stub |
@@ -0,0 +1,94 @@
---
name: index-related-work
description: Ingest a paper or repository into the related work index
---
# Index Related Work
## Purpose
Create a structured summary of a related paper, repository, or blog post and
add it to `.agents/memory/related-work/` for future reference. This builds the
agent's knowledge base for making informed decisions about training, evaluation,
and architecture choices.
## Prerequisites
- Access to the paper/repo (URL, PDF, or local clone).
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `source` | Yes | URL, citation, or local path |
| `type` | Yes | `paper`, `repo`, or `blog` |
| `tags` | No | List of tags (default: inferred from content) |
## Steps
### 1. Extract key information
For **papers**: Read abstract, method section, experimental setup, and results.
For **repos**: Read README, key source files, and training scripts.
For **blogs**: Read the full post.
Focus on:
- What problem does it solve?
- What architecture/technique is used?
- How does it relate to FastVideo's approach?
### 2. Create the index entry
Write to `.agents/memory/related-work/<slug>.md`:
```markdown
---
title: <title>
source: <URL or citation>
type: paper | repo | blog
date_indexed: <ISO-8601>
tags: [world-model, distillation, evaluation, ...]
---
## Summary
<1-2 paragraph summary.>
## Key Differences from FastVideo
- <comparison points>
## Actionable Insights
- <what we could adopt or adapt>
```
### 3. Update the catalog
If `.agents/memory/related-work/_catalog.md` exists, append the new entry.
If not, create it:
```markdown
# Related Work Catalog
| Slug | Title | Type | Tags | Date |
|------|-------|------|------|------|
| <slug> | <title> | <type> | <tags> | <date> |
```
## Outputs
- New file in `.agents/memory/related-work/<slug>.md`.
- Updated catalog.
## Example Usage
```
Index the Self-Forcing paper:
source: https://arxiv.org/abs/2406.xxxxx
type: paper
tags: [world-model, self-forcing, distillation]
```
## References
- `.agents/memory/related-work/README.md` — schema documentation
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
+7
View File
@@ -0,0 +1,7 @@
{"name": "launch-experiment", "description": "Generate and execute a training launch command for FastVideo models", "path": "launch-experiment/SKILL.md", "status": "draft", "trust": "low"}
{"name": "monitor-experiment", "description": "Poll a running W&B training run for progress and emit structured alerts", "path": "monitor-experiment/SKILL.md", "status": "draft", "trust": "low"}
{"name": "summarize-run", "description": "Extract a W&B run summary into a structured experiment report", "path": "summarize-run/SKILL.md", "status": "draft", "trust": "low"}
{"name": "log-experiment", "description": "Append or update an experiment entry in the experiment journal", "path": "log-experiment/SKILL.md", "status": "draft", "trust": "low"}
{"name": "evaluate-video-quality", "description": "Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)", "path": "evaluate-video-quality/SKILL.md", "status": "draft", "trust": "low"}
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
+127
View File
@@ -0,0 +1,127 @@
---
name: launch-experiment
description: Generate and execute a training launch command for FastVideo models
---
# Launch Experiment
## Purpose
Construct a fully-specified `torchrun` training command for a FastVideo model
given a target pipeline, dataset, and hyperparameter overrides. This skill
automates the boilerplate of setting environment variables, picking the right
entrypoint, and applying defaults from the closest example script.
## Prerequisites
- The repo is cloned and `fastvideo` is installed (`uv pip install -e .[dev]`).
- Dataset is preprocessed (see `docs/training/data_preprocess.md`).
- `WANDB_API_KEY` is set in the environment (or `WANDB_MODE=offline` for local).
- GPU resources are available (multi-GPU requires NCCL).
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `pipeline` | Yes | Training pipeline type: `finetune`, `distill-dmd`, `self-forcing`, `lora`, `consistency` |
| `model` | Yes | Model family: `wan-t2v-1.3B`, `wan-i2v-14B`, `ltx2`, `matrixgame` |
| `data_path` | Yes | Path to preprocessed dataset (parquet) |
| `num_gpus` | Yes | Number of GPUs |
| `overrides` | No | Dict of hyperparameter overrides (any CLI arg) |
| `output_dir` | No | Output directory (default: `outputs/<model>_<pipeline>`) |
| `run_name` | No | W&B run name (default: auto-generated) |
## Steps
### 1. Identify the training entrypoint
| Pipeline | Entrypoint |
|----------|-----------|
| `finetune` (Wan T2V) | `fastvideo/training/wan_training_pipeline.py` |
| `finetune` (Wan I2V) | `fastvideo/training/wan_i2v_training_pipeline.py` |
| `finetune` (LTX-2) | `fastvideo/training/ltx2_training_pipeline.py` |
| `finetune` (MatrixGame) | `fastvideo/training/matrixgame_training_pipeline.py` |
| `distill-dmd` | `fastvideo/training/wan_distillation_pipeline.py` |
| `self-forcing` | `fastvideo/training/wan_self_forcing_distillation_pipeline.py` |
### 2. Resolve default hyperparameters
Find the closest example script in `examples/training/` for the model:
| Model | Example Script Directory |
|-------|-------------------------|
| `wan-t2v-1.3B` | `examples/training/finetune/wan_t2v_1.3B/crush_smol/` |
| `wan-i2v-14B` | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
| `ltx2` | `examples/training/finetune/ltx2/` |
| `matrixgame` | `examples/training/finetune/MatrixGame2.0/` |
| `distill-dmd` | `scripts/distill/v1_distill_dmd_wan.sh` |
Read the script to extract default values for:
- `--learning_rate`, `--train_batch_size`, `--sp_size`, `--tp_size`
- `--num_latent_t`, `--num_height`, `--num_width`, `--num_frames`
- `--gradient_accumulation_steps`, `--max_train_steps`
- `--mixed_precision`, `--weight_decay`, `--max_grad_norm`
- `--validation_steps`, `--validation_sampling_steps`
### 3. Set environment variables
```bash
export WANDB_API_KEY="${WANDB_API_KEY}"
export WANDB_BASE_URL="https://api.wandb.ai"
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export TOKENIZERS_PARALLELISM=false
export TRITON_CACHE_DIR=/tmp/triton_cache
```
### 4. Construct the torchrun command
```bash
torchrun --nnodes 1 --nproc_per_node <num_gpus> \
<entrypoint> \
--pretrained_model_name_or_path <model_hf_id> \
--data_path "<data_path>" \
--output_dir "<output_dir>" \
--wandb_run_name "<run_name>" \
--tracker_project_name "<project_name>" \
--log_validation \
<...all hyperparameters...>
```
### 5. Log to experiment journal
After launching, append an entry to `.agents/memory/experiment-journal/README.md`:
```markdown
## [YYYY-MM-DD] Experiment: <run_name>
- **Hypothesis**: <user-provided or auto-generated>
- **Config**: model=<model>, lr=<lr>, sp_size=<sp>, gpus=<n>, script=<entrypoint>
- **W&B run**: <pending — will be updated by monitor skill>
- **Status**: running
```
## Outputs
- A ready-to-execute shell command.
- An experiment journal entry.
## Example Usage
```
Launch a Wan T2V 1.3B finetune on 4 GPUs with lr=5e-5 and max_train_steps=1000:
pipeline: finetune
model: wan-t2v-1.3B
data_path: data/crush_smol_preprocessed/
num_gpus: 4
overrides:
learning_rate: 5e-5
max_train_steps: 1000
```
## References
- `examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
- `scripts/distill/v1_distill_dmd_wan.sh`
- `docs/training/finetune.md` (training arguments table)
- `fastvideo/training/trackers.py` (tracker initialization)
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
+87
View File
@@ -0,0 +1,87 @@
---
name: log-experiment
description: Append or update an experiment entry in the experiment journal
---
# Log Experiment
## Purpose
Create or update an entry in `.agents/memory/experiment-journal/README.md` to maintain
a living record of all experiments and their outcomes.
## Prerequisites
- `.agents/memory/experiment-journal/README.md` exists.
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `name` | Yes | Experiment name / identifier |
| `hypothesis` | No | What you expected to learn |
| `config` | Yes | Key config: model, lr, sp_size, gpus, script |
| `wandb_run` | No | W&B run ID or URL |
| `duration` | No | Total wall time |
| `metrics` | No | Key metrics dict (loss, step_time, grad_norm) |
| `checkpoint` | No | Path to checkpoint |
| `insight` | No | What was learned |
| `status` | Yes | `running`, `completed`, `failed`, `abandoned` |
| `lessons` | No | Paths to related lesson files |
## Steps
### 1. Check for existing entry
Search `.agents/memory/experiment-journal/README.md` for an entry with the same name.
If found, update it instead of creating a duplicate.
### 2. Format the entry
```markdown
## [YYYY-MM-DD] Experiment: <name>
- **Hypothesis**: <hypothesis or "N/A">
- **Config**: model=<model>, lr=<lr>, sp_size=<sp>, gpus=<n>, script=<script>
- **W&B run**: <wandb_run or "pending">
- **Duration**: <duration or "in progress">
- **Key metrics**: loss=<loss>, step_time=<step_time>, grad_norm=<grad_norm>
- **Checkpoint**: <checkpoint or "N/A">
- **Insight**: <insight or "pending">
- **Status**: <status>
- **Related lessons**: <lessons or "none">
```
### 3. Insert at the top of the journal
New entries go at the top of the file (after the header), so the most recent
experiments are always visible first.
### 4. Warn on duplicates
If a similar experiment name exists with `status: completed`, warn that this
may be a repeat. If it's `status: running`, assume this is an update.
## Outputs
- Updated `.agents/memory/experiment-journal/README.md`.
## Example Usage
```
Log a completed experiment:
name: wan-t2v-finetune-lr5e5-sp4
config: model=wan-t2v-1.3B, lr=5e-5, sp_size=4, gpus=4
wandb_run: fastvideo/training/run_abc123
duration: 2h 15m
metrics: {loss: 0.065, step_time: 2.3, grad_norm: 0.35}
checkpoint: outputs/wan_finetune/checkpoint-1000
insight: LR 5e-5 converges 30% faster than 1e-5 with no quality loss
status: completed
```
## References
- `.agents/memory/experiment-journal/README.md` — journal file
- `.agents/workflows/experiment-lifecycle.md` — when to log
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
+134
View File
@@ -0,0 +1,134 @@
---
name: monitor-experiment
description: Poll a running W&B training run for progress and emit structured alerts
---
# Monitor Experiment
## Purpose
Continuously (or on-demand) check a running experiment's W&B metrics and emit
alerts for anomalies. Supports the "30-minute quality check" paradigm: after
the first 30 minutes of a long training run, produce a checkpoint quality
report before committing more resources.
## Prerequisites
- `WANDB_API_KEY` is set in the environment.
- The experiment is actively logging to W&B (not in `WANDB_MODE=offline`).
- For offline mode: read from local `wandb-summary.json` instead.
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `run_id` | Yes* | W&B run ID (e.g., `entity/project/run_id`) |
| `output_dir` | Yes* | Local output directory (for offline mode fallback) |
| `poll_interval` | No | Seconds between polls (default: 60) |
| `alert_on` | No | List of alert conditions to enable (default: all) |
\* One of `run_id` or `output_dir` is required.
## Steps
### 1. Connect to the run
**Online mode** (preferred):
```python
import wandb
api = wandb.Api()
run = api.run("<run_id>")
```
**Offline fallback**:
```python
import json
summary_path = f"{output_dir}/tracker/wandb/latest-run/files/wandb-summary.json"
with open(summary_path) as f:
summary = json.load(f)
```
### 2. Track key metrics
| Metric | W&B Key | Description |
|--------|---------|-------------|
| Training loss | `train_loss` | Primary training loss |
| Gradient norm | `grad_norm` | Gradient magnitude |
| Step time | `step_time` | Wall-clock seconds per step |
| Learning rate | `learning_rate` | Current LR |
| Avg step time | `avg_step_time` | Running average step time |
| Validation videos | `validation_videos_*` | Generated validation samples |
### 3. Evaluate alert conditions
| Alert | Condition | Severity |
|-------|-----------|----------|
| **Loss spike** | `current_loss > 3 × rolling_avg_loss` | 🔴 Critical |
| **NaN/Inf gradient** | `grad_norm` is NaN or Inf | 🔴 Critical |
| **Step time regression** | `step_time > 2 × baseline_step_time` | 🟡 Warning |
| **No progress** | No new W&B logs for > 10 minutes | 🟡 Warning |
| **Loss plateau** | Loss change < 1% over last 100 steps | 🟢 Info |
### 4. Emit structured status
Output format (agent-consumable):
```json
{
"run_id": "...",
"step": 500,
"metrics": {
"train_loss": 0.078,
"grad_norm": 0.41,
"step_time": 2.5,
"learning_rate": 1e-6
},
"alerts": [
{"type": "loss_spike", "severity": "critical", "message": "Loss jumped to 0.45 (avg: 0.08)"}
],
"status": "running"
}
```
### 5. 30-Minute Quality Check
After the first 30 minutes of wall-clock time:
1. Summarize the loss curve shape (decreasing? at what rate?).
2. Check if validation videos have been generated.
3. Report step count, loss at start vs. current, and estimated time to completion.
4. Produce a go/no-go recommendation.
```markdown
## 30-Minute Check: <run_name>
- **Steps completed**: 150
- **Loss**: 0.12 → 0.08 (↓ 33%)
- **Grad norm**: stable at ~0.4
- **Step time**: 2.5s/step (consistent)
- **Validation videos**: 5 generated at step 100
- **Recommendation**: ✅ Continue — loss is decreasing normally
```
## Outputs
- Structured JSON status updates.
- Alert messages for anomalous conditions.
- 30-minute checkpoint quality report.
## Example Usage
```
Monitor W&B run "fastvideo/Wan_distillation/abc123":
run_id: fastvideo/Wan_distillation/abc123
poll_interval: 120
alert_on: [loss_spike, nan_gradient, step_time_regression]
```
## References
- `fastvideo/training/trackers.py` — `WandbTracker` implementation
- `fastvideo/tests/training/Vanilla/test_training_loss.py` — how summaries are compared
- `fastvideo/tests/training/Vanilla/a40_reference_wandb_summary.json` — reference summary format
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
@@ -0,0 +1,82 @@
---
name: search-related-work
description: Query the related work index for relevant papers, repos, or comparisons
---
# Search Related Work
## Purpose
Search through `.agents/memory/related-work/` to find indexed papers, repos,
or blog posts relevant to a query. Use this when you need to understand how
other work compares to FastVideo's approach, or when looking for techniques
to adopt.
## Prerequisites
- The related work index has entries (`.agents/memory/related-work/*.md`).
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `query` | Yes | Natural language query |
| `tags` | No | Filter by tags (e.g., `[distillation, evaluation]`) |
| `type` | No | Filter by type (`paper`, `repo`, `blog`) |
## Steps
### 1. Search the index
Use grep-based search through `.agents/memory/related-work/`:
```bash
# Search by content
grep -rl "<query>" .agents/memory/related-work/
# Search by tags (in frontmatter)
grep -l "tags:.*<tag>" .agents/memory/related-work/*.md
```
### 2. Rank results
For each matching file:
1. Read the file.
2. Score relevance to the query based on:
- Title match
- Tag match
- Content match (summary, differences, insights)
3. Return top results.
### 3. Format output
```markdown
## Related Work Search: "<query>"
### 1. <Title> (relevance: high)
- **Source**: <URL>
- **Tags**: <tags>
- **Key insight**: <most relevant excerpt>
- **File**: `.agents/memory/related-work/<slug>.md`
### 2. <Title> (relevance: medium)
...
```
## Outputs
- Ranked list of relevant related work entries with excerpts.
## Example Usage
```
Search for work related to video quality evaluation metrics:
query: "video generation quality evaluation metrics"
tags: [evaluation]
```
## References
- `.agents/memory/related-work/README.md` — index schema
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
+137
View File
@@ -0,0 +1,137 @@
---
name: summarize-run
description: Extract a W&B run summary into a structured experiment report
---
# Summarize Run
## Purpose
After a training run completes (or at any checkpoint), extract key metrics from
the W&B run summary and produce a structured markdown report. Supports both
online (W&B API) and offline (local `wandb-summary.json`) modes.
## Prerequisites
- Run has completed or reached a checkpoint with a saved summary.
- For online: `WANDB_API_KEY` set in environment.
- For offline: access to `<output_dir>/tracker/wandb/latest-run/files/wandb-summary.json`.
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `run_id` | Yes* | W&B run ID for online access |
| `output_dir` | Yes* | Local output dir for offline access |
| `reference_run` | No | Path to reference `wandb-summary.json` for comparison |
| `experiment_name` | No | Name for the journal entry (default: from W&B) |
\* One of `run_id` or `output_dir` is required.
## Steps
### 1. Load run summary
**Online**:
```python
import wandb
api = wandb.Api()
run = api.run("<run_id>")
summary = dict(run.summary)
config = dict(run.config)
```
**Offline** (existing codebase pattern from `fastvideo/tests/training/`):
```python
import json
summary_path = f"{output_dir}/tracker/wandb/latest-run/files/wandb-summary.json"
with open(summary_path) as f:
summary = json.load(f)
```
### 2. Extract key fields
| Field | Source | Description |
|-------|--------|-------------|
| `train_loss` | `summary["train_loss"]` | Final training loss |
| `avg_step_time` | `summary["avg_step_time"]` | Average seconds per step |
| `step_time` | `summary["step_time"]` | Last step time |
| `grad_norm` | `summary["grad_norm"]` | Final gradient norm |
| `learning_rate` | `summary["learning_rate"]` | Final LR |
| `_step` | `summary["_step"]` | Total steps completed |
| `_runtime` | `summary["_runtime"]` | Total wall-clock seconds |
| `validation_videos_*` | `summary[key]` | Validation video artifacts |
### 3. Compare against reference (optional)
Follow the pattern in `fastvideo/tests/training/Vanilla/test_training_loss.py`:
```python
# Fields to compare
compare_fields = ["train_loss", "grad_norm", "avg_step_time"]
tolerance = 0.05 # 5% relative tolerance
for field in compare_fields:
ref_val = reference_summary[field]
cur_val = summary[field]
diff_pct = abs(cur_val - ref_val) / abs(ref_val) * 100
status = "✅" if diff_pct < tolerance * 100 else "⚠️"
print(f"{status} {field}: {cur_val:.4f} (ref: {ref_val:.4f}, diff: {diff_pct:.1f}%)")
```
### 4. Generate report
```markdown
# Run Summary: <experiment_name>
| Metric | Value | Reference | Diff |
|--------|-------|-----------|------|
| Train Loss | 0.0788 | 0.0800 | -1.5% ✅ |
| Avg Step Time | 2.81s | 2.80s | +0.4% ✅ |
| Grad Norm | 0.408 | 0.410 | -0.5% ✅ |
| Total Steps | 500 | — | — |
| Wall Time | 23m 30s | — | — |
## Configuration
- Model: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
- Learning Rate: 1e-6
- Batch Size: 1
- GPUs: 8 × (SP=1, TP=1)
- Mixed Precision: bf16
## Validation Videos
<list of validation video paths if available>
## Notes
<any observations or anomalies>
```
### 5. Update experiment journal
Append or update the experiment's entry in `.agents/memory/experiment-journal/README.md`
with the final metrics and status.
## Outputs
- Structured markdown report.
- Updated experiment journal entry.
## Example Usage
```
Summarize the run in output directory "outputs/wan_finetune":
output_dir: outputs/wan_finetune
reference_run: fastvideo/tests/training/Vanilla/a40_reference_wandb_summary.json
experiment_name: wan-t2v-finetune-lr1e6
```
## References
- `fastvideo/tests/training/Vanilla/test_training_loss.py` — reference comparison pattern
- `fastvideo/tests/training/Vanilla/a40_reference_wandb_summary.json` — example summary
- `fastvideo/tests/training/lora/test_lora_training.py` — LoRA summary comparison
- `fastvideo/training/trackers.py` — tracker summary generation
## Changelog
| Date | Change |
|------|--------|
| 2026-03-02 | Initial version |
@@ -0,0 +1,54 @@
---
description: How to develop, validate, and register a new evaluation metric
---
# Evaluation Development SOP
Standard procedure for adding new video quality evaluation metrics to the
FastVideo agent toolkit.
## When to Use
- You need a metric that doesn't exist in `.agents/memory/evaluation-registry/README.md`.
- An existing metric needs significant changes to its methodology.
- You're exploring a new evaluation approach.
## Steps
### 1. Research
- Search `.agents/memory/related-work/` for existing evaluation approaches.
- Check the `evaluation_registry.md` for current metrics and their limitations.
- Review literature: FVD, CLIP-Score, human preference, etc.
### 2. Prototype
- Write a standalone script in `.agents/exploration/<metric-name>.md`.
- Keep it simple: one script, minimal dependencies.
- Test on a few known-good and known-bad video samples.
### 3. Validate
- **Known-good test**: Metric should score high on reference-quality videos.
- **Known-bad test**: Metric should score low on degraded/unrelated videos.
- **Sensitivity test**: Small quality differences should produce meaningful
score differences.
- Document thresholds and their justification.
### 4. Register
Update `.agents/memory/evaluation-registry/README.md`:
- Add the metric with status `Active`.
- Document location, thresholds, and trust level.
### 5. Integrate
Update `.agents/skills/evaluate-video-quality.md`:
- Add the new metric as a section.
- Include code examples and interpretation guide.
### 6. Document
- Move the exploration log content into the skill.
- Clean up the exploration file or mark it as `promoted`.
- If anything went wrong during development, create a lesson.
@@ -0,0 +1,47 @@
---
description: When and how to log experiments in the experiment journal
---
# Experiment Journaling SOP
Ensures every experiment is properly recorded with context and outcomes.
## When to Log
**Always.** Every experiment — even quick tests — should be journaled.
## Steps
### 1. Before Launch — Create Draft Entry
Use the `log-experiment` skill with `status: running`:
- Include hypothesis and config.
- Leave metrics, duration, and insight blank.
### 2. After 30-Minute Check — Update with Initial Metrics
Update the entry with:
- Current loss and its trajectory direction.
- Step time.
- Number of validation videos generated.
- Preliminary go/no-go assessment.
### 3. On Completion — Fill Final Entry
Update the entry with `status: completed`:
- Final loss, grad norm, avg step time.
- Total duration and steps.
- Checkpoint path.
- Key insight.
### 4. On Failure — Document Failure Mode
Update the entry with `status: failed`:
- What went wrong (OOM, NaN, crash, etc.).
- At what step the failure occurred.
- Create a lesson in `.agents/lessons/` for non-trivial failures.
### 5. Cross-Reference
- Link related lessons: `**Related lessons**: .agents/lessons/<filename>.md`
- Link related experiments: if this is a follow-up, reference the prior entry.
+87
View File
@@ -0,0 +1,87 @@
---
description: End-to-end experiment lifecycle from hypothesis to lessons learned
---
# Experiment Lifecycle SOP
Standard operating procedure for running ML training experiments on
FastVideo-WorldModel. Every experiment should follow this flow.
## Overview
```
Plan → Launch → Monitor → Summarize → Journal → Reflect
```
## Steps
### 1. Plan the Experiment
Before launching:
- [ ] Define a clear **hypothesis** (what you expect to learn).
- [ ] Select the **model** and **pipeline** type (finetune, distill, lora, etc.).
- [ ] Prepare the **dataset** (preprocessed into parquet format).
- [ ] Review existing experiments in `.agents/memory/experiment-journal/README.md` for related work.
- [ ] Check `.agents/lessons/` for known pitfalls with this configuration.
- [ ] Document the plan in the experiment journal as a draft entry.
### 2. Launch the Experiment
Use the `launch-experiment` skill:
- Provide: pipeline, model, data_path, num_gpus, and any hyperparameter overrides.
- The skill generates the `torchrun` command and creates a journal entry.
- Verify the command looks correct before executing.
Reference: `.agents/skills/launch-experiment.md`
### 3. Monitor the Experiment
Use the `monitor-experiment` skill:
- Provide the W&B run ID (or output_dir for offline).
- Monitor alerts: loss spikes, NaN gradients, step time regressions.
- At the **30-minute mark**: perform the quality check.
- Is loss decreasing?
- Are validation videos reasonable?
- Is step time consistent?
- **Decision point**: Continue or abort based on the 30-min check.
Reference: `.agents/skills/monitor-experiment.md`
### 4. Summarize the Run
After completion (or at any checkpoint), use the `summarize-run` skill:
- Extract final metrics from W&B summary.
- Compare against reference runs if available.
- Generate a structured report.
Reference: `.agents/skills/summarize-run.md`
### 5. Update the Experiment Journal
Use the `log-experiment` skill to update the journal entry:
- Fill in final metrics, duration, checkpoint paths.
- Record the key insight learned.
- Set status to `completed`, `failed`, or `abandoned`.
Reference: `.agents/skills/log-experiment.md`
### 6. Reflect and Capture Lessons
After every experiment:
- **What went right?** → Note in the journal insight field.
- **What went wrong?** → Create a lesson in `.agents/lessons/`:
- Use the template in `.agents/lessons/README.md`.
- Cross-reference the experiment journal entry.
- **What was surprising?** → Consider creating an exploration log if this
warrants further investigation.
Reference: `.agents/workflows/lesson-capture.md`
## Validation Criteria
This SOP is validated when an agent can:
1. Follow steps 1–6 end-to-end for a minimal training run
(e.g., `examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
with `--max_train_steps 5`).
2. Produce a complete experiment journal entry.
3. Generate a run summary report.
+71
View File
@@ -0,0 +1,71 @@
---
description: Post-experiment reflection to capture lessons learned
---
# Lesson Capture SOP
Systematic procedure for turning experiment outcomes into persistent knowledge.
## When to Use
After **every** completed or failed experiment. Even successful experiments
can yield lessons (e.g., "LR 5e-5 works better than 1e-5 for LoRA").
## Steps
### 1. Review the Experiment
Read the experiment journal entry. Ask:
- Did anything go wrong?
- Was anything surprising?
- Did anything take longer than expected?
- Was a workaround needed?
### 2. Decide: Lesson or Not?
| Situation | Action |
|-----------|--------|
| Something broke | Create a lesson (category: `infrastructure` or `data`) |
| Hyperparameter choice mattered | Create a lesson (category: `hyperparameter`) |
| Porting issue found | Create a lesson (category: `porting`) |
| Evaluation metric was misleading | Create a lesson (category: `evaluation`) |
| Everything went smoothly | No lesson needed, but note in the journal insight |
### 3. Create the Lesson File
In `.agents/lessons/`, create `<YYYY-MM-DD>_<short-slug>.md`:
```markdown
---
date: <ISO-8601>
experiment: <journal entry reference>
category: hyperparameter | data | infrastructure | evaluation | porting
severity: critical | important | minor
---
# <Short Descriptive Title>
## What Happened
<description>
## Root Cause
<analysis>
## Fix / Workaround
<resolution>
## Prevention
<how to avoid in future>
```
### 4. Cross-Reference
- Update the experiment journal entry with a link to the lesson file.
- If a similar lesson already exists, add a reference or update it.
### 5. Periodic Pattern Review
Every ~10 lessons, scan for patterns:
- Multiple lessons in the same category → consider a new skill or SOP.
- Repeated mistakes → strengthen the relevant SOP with a checklist item.
- Infrastructure issues → propose a codebase fix.
+67
View File
@@ -0,0 +1,67 @@
---
description: Synchronize the STATUS.md dashboard by scanning .agents/ directories
---
# Sync Dashboard
Updates `.agents/STATUS.md` by scanning the skills, workflows, memory, lessons,
and exploration directories to reflect what actually exists on disk.
## When to Use
- After adding, removing, or renaming any file in `.agents/`.
- Periodically (e.g., at end of each conversation session).
- When the dashboard feels out of date.
## Steps
### 1. Scan directories
List all files in each directory:
```bash
echo "=== Skills ==="
ls -1 .agents/skills/*.md 2>/dev/null | grep -v SKILL_TEMPLATE
echo "=== Workflows ==="
ls -1 .agents/workflows/*.md 2>/dev/null
echo "=== Memory ==="
ls -1 .agents/memory/*.md 2>/dev/null
ls -1 .agents/memory/related-work/*.md 2>/dev/null | grep -v README
echo "=== Lessons ==="
ls -1 .agents/lessons/*.md 2>/dev/null | grep -v README
echo "=== Exploration ==="
ls -1 .agents/exploration/*.md 2>/dev/null | grep -v README
```
### 2. Compare with STATUS.md
For each file found:
- If it's in STATUS.md → leave it (preserve status/trust/tested fields).
- If it's NOT in STATUS.md → add it with status `🔴 Stub`, trust `None`, tested `❌`.
For each entry in STATUS.md:
- If the file no longer exists → mark it as `❌ Removed` or delete the row.
### 3. Update counts
Recalculate the summary table at the top:
- Count files per category.
- Count by status (Ready, Draft, Stub).
### 4. Update timestamp
Set `_Last synced: <current date>_` at the top of STATUS.md.
### 5. Review
Read through the updated STATUS.md for accuracy. Flag anything that looks wrong.
## Notes
- Do NOT change trust levels during sync — those are set manually after testing.
- Do NOT change status during sync — status changes require actual validation.
- This workflow only handles structural sync (file existence), not content review.
@@ -0,0 +1,46 @@
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"description": "Wan2.1 T2V 1.3B inference performance",
"model": {
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"model_short_name": "Wan2.1-T2V-1.3B"
},
"init_kwargs": {
"num_gpus": 2,
"flow_shift": 7.0,
"sp_size": 2,
"tp_size": 1,
"vae_sp": true,
"vae_tiling": true,
"text_encoder_precisions": ["fp32"]
},
"generation_kwargs": {
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 4,
"guidance_scale": 3,
"embedded_cfg_scale": 6,
"seed": 1024,
"fps": 24,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
},
"test_prompts": [
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
],
"run_config": {
"num_warmup_runs": 1,
"num_measurement_runs": 3,
"required_gpus": 2
},
"thresholds": {
"L40S": {
"max_generation_time_s": 34.0,
"max_peak_memory_mb": 11000.0
},
"default": {
"max_generation_time_s": 120.0,
"max_peak_memory_mb": 30000.0
}
}
}
+31 -12
View File
@@ -139,18 +139,6 @@ steps:
- TEST_TYPE=training_vsa
agents:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "pyproject.toml"
@@ -183,6 +171,37 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
- path:
- "fastvideo/models/dits/**"
- "fastvideo/pipelines/**"
- "fastvideo/attention/**"
- "fastvideo/layers/**"
- "fastvideo/worker/**"
- "fastvideo/entrypoints/**"
- "fastvideo/tests/performance/**"
- ".buildkite/performance-benchmarks/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Performance Tests"
env:
- TEST_TYPE=performance
agents:
queue: "default"
- path:
- "fastvideo/entrypoints/openai/**"
- "fastvideo/entrypoints/cli/serve.py"
- "fastvideo/tests/entrypoints/test_openai_api_integration.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "API Server Tests"
env:
- TEST_TYPE=api_server
agents:
queue: "default"
# - path:
# - "scripts/lora_extraction/**"
# - "pyproject.toml"
+10 -5
View File
@@ -51,6 +51,7 @@ else
fi
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
MODAL_SSIM_TEST_FILE="fastvideo/tests/modal/ssim_test.py"
if [ -z "${TEST_TYPE:-}" ]; then
log "Error: TEST_TYPE environment variable is not set"
@@ -75,7 +76,7 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
@@ -89,10 +90,6 @@ case "$TEST_TYPE" in
log "Running training VSA tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
;;
"inference_sta")
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"kernel_tests")
log "Running kernel tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
@@ -122,6 +119,14 @@ case "$TEST_TYPE" in
log "Running LoRA extraction tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
;;
"performance")
log "Running performance tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_performance_tests"
;;
"api_server")
log "Running API server integration tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_api_server_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+43
View File
@@ -0,0 +1,43 @@
## Purpose
<!-- What does this PR do? Link the related issue if applicable. -->
Fixes #
## Changes
<!-- Describe your changes concisely. What approach did you take? -->
-
## Test Plan
<!-- How did you verify your changes? Paste exact commands and output. -->
```bash
# Commands you ran
```
## Test Results
<!-- Paste test output, before/after comparisons, or SSIM scores for model changes. -->
<details>
<summary>Test output</summary>
```
# Paste output here
```
</details>
## Checklist
- [ ] I ran `pre-commit run --all-files` and fixed all issues
- [ ] I added or updated tests for my changes
- [ ] I updated documentation if needed
- [ ] I considered GPU memory impact of my changes
**For model/pipeline changes, also check:**
- [ ] I verified SSIM regression tests pass
- [ ] I updated the support matrix if adding a new model
+2 -65
View File
@@ -47,16 +47,6 @@ on:
required: false
default: false
type: boolean
run_inference_test_STA:
description: "Run inference-test-STA"
required: false
default: false
type: boolean
run_precision_test_STA:
description: "Run precision-test-STA"
required: false
default: false
type: boolean
run_precision_test_VSA:
description: "Run precision-test-VSA"
required: false
@@ -90,8 +80,6 @@ jobs:
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
unit-test: ${{ steps.filter.outputs.unit-test }}
steps:
@@ -106,12 +94,6 @@ jobs:
- 'docker/Dockerfile.python3.10'
- 'docker/Dockerfile.python3.11'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/sliding_tile_attn/**'
- 'csrc/attn/sliding_tile_attn/tk/**'
- 'csrc/attn/sliding_tile_attn/setup.py'
- 'csrc/attn/sliding_tile_attn/config_sta.py'
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/video_sparse_attn/**'
- 'csrc/attn/video_sparse_attn/tk/**'
@@ -148,13 +130,6 @@ jobs:
- 'fastvideo/**'
- *common-paths
- *vsa-kernel-paths
inference-test-STA:
- 'fastvideo/**'
- *common-paths
- *sta-kernel-paths
precision-test-STA:
- *common-paths
- *sta-kernel-paths
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
@@ -282,44 +257,6 @@ jobs:
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
inference-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "inference-test-STA"
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: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-STA"
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_sta.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-VSA:
needs: change-filter
if: >-
@@ -378,7 +315,7 @@ jobs:
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, inference-test-STA, precision-test-STA, precision-test-VSA]
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:
@@ -395,7 +332,7 @@ jobs:
- 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", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
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
-249
View File
@@ -1,249 +0,0 @@
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/attn/sliding_tile_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/sliding_tile_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']
torch-version: ['2.5.1', '2.6.0']
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
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.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.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.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-version }}+cu${{ matrix.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
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==${{ matrix.torch-version }} --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 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/sliding_tile_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/sliding_tile_attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.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/sliding_tile_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/sliding_tile_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/sliding_tile_attn/dist/
+4
View File
@@ -83,3 +83,7 @@ docs/distillation/examples/
dmd_t2v_output/
preprocess_output_text/
.claude/
.codex/
openspec/
+1 -1
View File
@@ -60,7 +60,7 @@ repos:
hooks:
- id: mypy
args: [--python-version, '3.10', --follow-imports, "skip", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
additional_dependencies: [types-aiofiles, types-cachetools, types-setuptools, types-PyYAML, types-requests]
- repo: local
hooks:
- id: check-filenames
+14
View File
@@ -40,3 +40,17 @@
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
- linked issue/PR context,
- screenshots or sample outputs for UI/demo/docs changes.
## Agent Infrastructure
This repository is agent-friendly. Before doing any work, read:
1. `.agents/onboarding/README.md` — full onboarding guide with step-by-step instructions.
2. `.agents/memory/codebase-map/README.md` — structural index of the entire repository.
3. `.agents/skills/` — available agent skills (check if one exists before writing code).
4. `.agents/workflows/` — SOPs for common procedures (experiment lifecycle, evaluation, etc.).
5. `.agents/lessons/` — known pitfalls and their documented fixes.
If you are exploring a new procedure that has no existing SOP, document your
progress in `.agents/exploration/` and flag it for review at the end of your
session.
+5 -6
View File
@@ -44,15 +44,15 @@ FastVideo has the following features:
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
We recommend using [uv](https://docs.astral.sh/uv/) to create a clean environment. If you previously used Conda, switching to uv generally gives faster and more stable installs.
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Create and activate a new uv environment
uv venv --python 3.12 --seed
source .venv/bin/activate
# Install FastVideo
pip install fastvideo
uv pip install fastvideo
```
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
@@ -93,7 +93,6 @@ def main():
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+5 -5
View File
@@ -9,11 +9,12 @@ import datetime
import locale
import os
import re
import shutil
import subprocess
import sys
# Unlike the rest of the PyTorch this file must be python2 compliant.
# This script outputs relevant system environment info
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
# This script outputs relevant system environment info.
# Run it with: python collect_env.py
# Requires Python 3.10+ (matches fastvideo); uses shutil.which (Python 3.3+).
from collections import namedtuple
from fastvideo.envs import environment_variables
@@ -495,8 +496,7 @@ def get_pip_packages(run_lambda, patterns=None):
if pip_available:
cmd = [sys.executable, '-mpip', 'list', '--format=freeze']
elif os.environ.get("UV") is not None:
print("uv is set")
elif shutil.which("uv") is not None:
cmd = ["uv", "pip", "list", "--format=freeze"]
else:
raise RuntimeError(
+1 -7
View File
@@ -443,7 +443,6 @@
1025,
"fixed",
24,
-99999,
-99999
],
"auto_widget_states": {
@@ -491,11 +490,6 @@
"isAuto": true,
"value": -99999,
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
@@ -642,4 +636,4 @@
"VHS_KeepIntermediate": true
},
"version": 0.4
}
}
@@ -353,8 +353,7 @@
1024,
"fixed",
24,
"X://insert/path/here.mp4",
true
"X://insert/path/here.mp4"
],
"auto_widget_states": {
"height": {
@@ -401,11 +400,6 @@
"isAuto": true,
"value": "X://insert/path/here.mp4",
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
@@ -694,4 +688,4 @@
"VHS_KeepIntermediate": true
},
"version": 0.4
}
}
@@ -31,9 +31,6 @@ class InferenceArgs:
"image_path": ("STRING", {
"default": "X://insert/path/here.mp4"
}),
"enable_teacache": ([True, False], {
"default": False
}),
}
}
@@ -57,7 +54,6 @@ class InferenceArgs:
seed,
fps,
image_path,
enable_teacache,
):
raw_args = {
"height": height,
@@ -69,7 +65,6 @@ class InferenceArgs:
"seed": seed,
"fps": fps,
"image_path": image_path,
"enable_teacache": enable_teacache,
}
# Filter out keys where value is -99999, handling different types properly
+1 -1
View File
@@ -552,7 +552,7 @@ app.registerExtension({
]
const floatWidgetNames = ["embedded_cfg_scale", "guidance_scale"]
const comboWidgetNames = ["vae_tiling", "vae_precision", "vae_sp", "text_encoder_precision", "precision",
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload"
]
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
+2 -1
View File
@@ -54,7 +54,8 @@ To use this:
1. **Set Context**: In your pipeline or generation loop, use the `set_forward_context` context manager.
2. **Access Context**: Inside your attention backend, use `get_forward_context()`.
See [`docs/attention/sta/index.md`](../sta/index.md) (Sliding Tile Attention) for an example of how complex configuration (window sizes) is passed this way.
See [`docs/attention/sta/index.md`](../sta/index.md) for a legacy STA example
of passing complex configuration (window sizes) through `ForwardContext`.
## 3. Adding Compiled Kernels (C++/CUDA)
+5 -2
View File
@@ -5,13 +5,16 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
## Supported Kernels
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
* **[Sliding Tile Attention (STA)](sta/index.md)**: STA kernel support is kept in
`fastvideo-kernel`; full FastVideo STA pipeline workflow is archived in
`sta_do_not_delete`.
* **Backend development guide**: See the developer guide at
[Attention Backend Development](../contributing/attention_backend.md).
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from source, which includes both STA and VSA kernels.
These instructions apply to building the `fastvideo-kernel` package from
source, which includes both STA and VSA kernels.
### Prerequisites
+71 -14
View File
@@ -1,27 +1,84 @@
# Sliding Tile Attention (STA)
Optimized attention for window-based video generation (e.g., HunyuanVideo).
STA inference integration is archived from `main`.
## Installation
The full STA pipeline code (including mask search and STA inference wiring in
`fastvideo/`) is preserved in:
STA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
- https://github.com/hao-ai-lab/FastVideo/tree/sta_do_not_delete
## Usage
In this branch, STA kernels in `fastvideo-kernel` are still kept.
```python
from fastvideo_kernel import sliding_tile_attention
## Why STA is not in `main`
# q, k, v: [batch_size, num_heads, seq_length, head_dim]
# window_size: List of (t, h, w) tiles. Tile size is (6, 8, 8).
# text_length: Number of text tokens (0-256)
We do not keep STA pipeline integration in `main` because we believe Video
Sparse Attention (VSA) is strictly better than STA for the actively maintained
FastVideo inference path.
out = sliding_tile_attention(
q, k, v,
window_size=[(3, 3, 3)], # Example window
text_length=256
)
## What to checkout for STA workflows
To run the full STA workflow, switch to the archived branch:
```bash
git fetch origin
git checkout sta_do_not_delete
```
## Mask Search (archive branch)
The reference script is:
- `examples/inference/sta_mask_search/inference_wan_sta.sh`
It runs two stages:
1. `STA_searching` (full search), output at
`inference_results/sta/mask_search_full`
2. `STA_tuning` (sparse tuning), output at
`inference_results/sta/mask_search_sparse`
Run:
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
bash examples/inference/sta_mask_search/inference_wan_sta.sh
```
## STA Inference (archive branch)
With a selected mask strategy, run inference with:
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
fastvideo generate \
--model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \
--num-gpus 2 \
--tp-size 2 \
--sp-size 2 \
--height 768 \
--width 1280 \
--num-frames 69 \
--num-inference-steps 50 \
--prompt "A cinematic wildlife shot of a lion walking in golden grasslands." \
--output-path outputs_video/STA/
```
Python usage on the archive branch can also set `STA_mode` in
`VideoGenerator.from_pretrained(...)`:
- `STA_searching`
- `STA_tuning`
- `STA_inference`
## Kernel-level API (current branch)
STA kernels remain available from `fastvideo-kernel`. See
[Attention overview](../index.md) for build instructions.
## Citation
If you use Sliding Tile Attention in your research, please cite:
+18 -17
View File
@@ -14,24 +14,11 @@ improving performance, or fixing a bug.
For a full install checklist, see `docs/getting_started/installation/gpu.md`.
## Local development (Conda + editable install)
## Local development (UV + editable install)
Install Miniconda:
If you previously used Conda for local setup, switch to uv for a faster and more stable development environment.
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
Create and activate a Conda environment:
```bash
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
Install `uv`:
```bash
curl -LsSf https://astral.sh/uv/install.sh | sh
@@ -39,6 +26,20 @@ curl -LsSf https://astral.sh/uv/install.sh | sh
wget -qO- https://astral.sh/uv/install.sh | sh
```
Create and activate a uv environment (recommended):
```bash
uv venv --python 3.12 --seed
source .venv/bin/activate
```
Conda alternative (supported):
```bash
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Clone the repo:
```bash
@@ -51,7 +52,7 @@ Install FastVideo in editable mode and set up hooks:
uv pip install -e .[dev]
# Optional: FlashAttention (builds native kernels)
uv pip install flash-attn --no-build-isolation
uv pip install flash-attn --no-build-isolation -v
# Linting, formatting, static typing
pre-commit install --hook-type pre-commit --hook-type commit-msg
+1 -1
View File
@@ -8,7 +8,7 @@ This guide explains how to add and run tests in FastVideo. The testing suite is
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., VSA, V-MoBA).
For now, we will focus on **SSIM Tests**.
+369
View File
@@ -0,0 +1,369 @@
# Training Architecture
!!! warning "Work in Progress"
This training architecture (`fastvideo/train/`) is under active development
and is replacing the older `fastvideo/training/` module. APIs, config
formats, and supported methods may change. See the
[Current Status](#current-status) section for what is implemented so far.
FastVideo's training framework (`fastvideo/train/`) is built around a
**pluggable, YAML-driven architecture** that cleanly separates **models**,
**training methods**, and **infrastructure** into independent, composable
layers. A single YAML config file is all that is needed to train any supported
model with any supported algorithm — no code changes required to mix and match.
---
## Motivation
Training video diffusion models involves a tangle of concerns: model loading,
noise scheduling, distillation algorithms, distributed strategies,
checkpointing, and validation. Existing training scripts tend to hard-wire
these together, making it painful to:
1. **Try a new distillation algorithm** on an existing model (requires forking
the training loop).
2. **Add a new model** to an existing algorithm (requires re-implementing
boilerplate).
3. **Switch distributed strategies** (FSDP, TP, SP) without touching algorithm
code.
4. **Resume, checkpoint, and validate** uniformly across all combinations.
The training framework solves this by making each axis of variation an
independent plugin.
---
## Architecture Overview
```
YAML Config
|
v
+------------------+ +---------------------+ +------------------+
| Models Layer | | Methods Layer | | Infrastructure |
| (per-role) | | (algorithm) | | Layer |
| | | | | |
| - ModelBase |<----| - TrainingMethod |---->| - Trainer |
| - CausalModelBase| | - single_train_step| | - Callbacks |
| | | - backward | | - Checkpoint |
| Roles: | | - optimizers | | - Tracker (W&B) |
| student | | | | - Dataloader |
| teacher | | Algorithms: | | |
| critic | | DMD2, SelfForcing, | | Distributed: |
| | | SFT, DFSFT | | HSDP, TP, SP |
+------------------+ +---------------------+ +------------------+
```
### Three Layers
| Layer | Responsibility | Extension point |
|-------|---------------|-----------------|
| **Models** (`fastvideo/train/models/`) | Load transformer + scheduler, define `predict_noise`, `predict_x0`, `add_noise`, `backward`. Each training role (student/teacher/critic) is an independent instance. | Subclass `ModelBase` (or `CausalModelBase` for streaming). |
| **Methods** (`fastvideo/train/methods/`) | Implement the training algorithm: own role models, define `single_train_step` + `backward`, manage optimizers/schedulers. | Subclass `TrainingMethod`. |
| **Infrastructure** (`fastvideo/train/trainer.py`, `utils/`, `callbacks/`) | Training loop, gradient accumulation, distributed setup, checkpointing (DCP), W&B tracking, validation, EMA, grad clipping. | Add callbacks; everything else is shared. |
---
## YAML-Driven Configuration
Everything is configured declaratively. The `_target_` field selects the Python
class to instantiate:
```yaml
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
rollout_mode: simulate
dmd_denoising_steps: [1000, 850, 700, 550, 350, 275, 200, 125]
generator_update_interval: 5
real_score_guidance_scale: 3.5
# ...
training:
distributed: { num_gpus: 8, sp_size: 1, tp_size: 1 }
data: { data_path: ..., num_latent_t: 20, num_frames: 77 }
optimizer: { learning_rate: 2.0e-6, betas: [0.0, 0.999] }
loop: { max_train_steps: 4000 }
checkpoint: { output_dir: outputs/my_run }
callbacks:
grad_clip: { max_grad_norm: 1.0 }
validation: { pipeline_target: ..., every_steps: 100 }
```
To switch from DMD2 to SFT, change the `method._target_` and remove the
teacher/critic — no code changes needed.
---
## Model Abstraction
### `ModelBase` — Standard (Bidirectional) Models
Every role gets its own `ModelBase` instance owning a `transformer` and
`noise_scheduler`. The base class defines:
- **`prepare_batch()`** — Convert raw dataloader output into forward-ready
`TrainingBatch`.
- **`add_noise()`** — Apply forward-process noise at a given timestep.
- **`predict_noise()` / `predict_x0()`** — Run the transformer and return
predictions.
- **`backward()`** — Backward pass that restores forward context (attention
metadata, timesteps).
- **`init_preprocessors()`** — Lazy-load VAE, build dataloader (called only on
the student).
### `CausalModelBase` — Streaming / Causal Models
Extends `ModelBase` with streaming inference primitives for causal video
generation:
```python
class CausalModelBase(ModelBase):
def clear_caches(self, *, cache_tag: str = "pos") -> None: ...
def predict_noise_streaming(
self, ..., cache_tag, store_kv, cur_start_frame
) -> Tensor | None: ...
def predict_x0_streaming(
self, ..., cache_tag, store_kv, cur_start_frame
) -> Tensor | None: ...
```
KV caches are **internal** to the model instance, keyed by `cache_tag`. The
method controls when to store (`store_kv=True`) vs. read-only
(`store_kv=False`), enabling block-by-block causal rollout during training.
---
## Training Methods
### DMD2 (Distribution Matching Distillation)
**Roles:** student (trainable) + teacher (frozen) + critic (trainable)
The student learns to generate clean video in few steps by matching the
teacher's score function, with a critic network providing a learned fake-score
baseline.
- **Rollout modes:**
- `simulate` — Student starts from pure noise and iteratively denoises
through the full step schedule.
- `data_latent` — Student denoises from a single randomly-noised data
sample.
- **Losses:** Generator loss (DMD gradient) + critic flow-matching loss, with
alternating updates (`generator_update_interval`).
### Self-Forcing (Causal DMD)
**Roles:** student (causal, trainable) + teacher (frozen) + critic (trainable)
Extends DMD2 for **streaming/causal video generation**. The key idea: during
training, the student processes video in temporal chunks, using its own
previously-denoised outputs as context for future chunks — simulating online
autoregressive rollout.
- Video is split into blocks of `chunk_size` latent frames.
- Each block is denoised through the student's step schedule; a random
early-exit step is sampled per block.
- After denoising a block, its output is fed back (with optional
`context_noise`) as KV cache context for subsequent blocks via
`predict_noise_streaming(store_kv=True)`.
- Supports SDE and ODE sampling during rollout.
- Selective gradient control: `enable_gradient_in_rollout`,
`start_gradient_frame`.
### Supervised Fine-Tuning (SFT)
**Roles:** student only
Standard flow-matching loss between predicted and ground-truth noise/x0.
### Diffusion-Forcing SFT (DFSFT)
**Roles:** student only
SFT with **inhomogeneous (per-chunk) timesteps** — each temporal chunk in a
video gets a different noise level. This trains the model to handle mixed-noise
inputs, which is a prerequisite for causal/streaming inference where earlier
frames are cleaner than later ones.
---
## Training Loop
The `Trainer` runs a standard loop with pluggable method and callbacks:
```
for step in range(start_step, max_steps):
for accum_iter in range(grad_accum_steps):
batch <- dataloader
loss_map, outputs, metrics <- method.single_train_step(batch, step)
method.backward(loss_map, outputs)
callbacks.on_before_optimizer_step() # grad clipping
method.optimizers_schedulers_step()
method.optimizers_zero_grad()
callbacks.on_training_step_end() # logging
checkpoint_manager.maybe_save(step)
callbacks.on_validation_begin() # periodic inference
```
### Callbacks
- **GradNormClipCallback** — Per-module gradient norm logging + global
clipping.
- **ValidationCallback** — Periodic inference sampling with configurable
pipeline, sampling steps, and guidance scale.
- **EMACallback** — Exponential moving average of student weights.
### Checkpointing
- DCP (Distributed Checkpoint) format, compatible with FSDP/HSDP.
- Saves: model weights, optimizer states, scheduler states, RNG states (per
role).
- Full resume support: auto-restores step counter and all RNG states.
---
## Getting Started
```bash
# Install
uv pip install -e .[dev]
# Run DMD2 distillation on Wan 2.1
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
--config examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml
# Run SFT fine-tuning
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
--config examples/train/finetune_wan2.1_t2v_1.3B_vsa_phase3.4_0.9sparsity.yaml
```
Example configs are in `examples/train/`.
---
## File Structure
```
fastvideo/train/
trainer.py # Training loop
models/
base.py # ModelBase, CausalModelBase ABCs
wan/wan.py # Wan 2.1 T2V model plugin
wangame/wangame.py # WanGame 2.1 I2V model plugin
wangame/wangame_causal.py # WanGame causal (streaming) plugin
methods/
base.py # TrainingMethod ABC
distribution_matching/
dmd2.py # DMD2 distillation
self_forcing.py # Self-Forcing (causal DMD)
fine_tuning/
finetune.py # Supervised fine-tuning
dfsft.py # Diffusion-forcing SFT
callbacks/
grad_clip.py # Gradient clipping + norm logging
validation.py # Periodic inference validation
ema.py # EMA weight averaging
entrypoint/
train.py # CLI entrypoint (torchrun)
utils/
config.py # YAML parser -> RunConfig
builder.py # build_from_config: model/method instantiation
training_config.py # TrainingConfig dataclass
dataloader.py # Dataset/dataloader construction
optimizer.py # Optimizer/scheduler construction
checkpoint.py # DCP save/resume
tracking.py # W&B tracker
```
---
## Current Status
| Component | Status |
|-----------|--------|
| Core framework (trainer, config, callbacks) | Implemented and tested |
| `WanModel` (Wan 2.1 T2V) | Implemented and tested |
| `WanGameModel` (WanGame 2.1 I2V) | Implemented and tested |
| `WanGameCausalModel` (streaming) | Implemented and tested |
| `WanCausalModel` (Wan T2V causal) | In progress |
| DMD2 method | Implemented and tested |
| Self-Forcing method | Implemented and tested |
| SFT method | Implemented and tested |
| DFSFT method | Implemented and tested |
| DCP checkpointing + resume | Implemented and tested |
| EMA callback | Implemented |
| Validation callback | Implemented and tested |
| Causal DMD inference pipeline | Implemented |
---
## Open Questions
We welcome community feedback on the following topics:
### Model Plugin API
The current `ModelBase` interface requires implementing 6 methods. Is this the
right granularity?
- Should `prepare_batch` be split into separate concerns (noise sampling,
timestep sampling, attention metadata)?
- Should `backward` be lifted out of the model and into the method/trainer?
### Causal Streaming Interface
`CausalModelBase` adds `predict_noise_streaming` / `predict_x0_streaming` with
cache management. Alternatives considered:
- **(a) Current:** Cache is internal to the model, keyed by `cache_tag`.
Simple but couples cache lifecycle to model.
- **(b) External cache:** Method owns the cache dict, passes it into predict
calls. More explicit but verbose.
- **(c) Context manager:** `with model.streaming_context(tag) as ctx: ...` —
cleaner lifecycle but harder to compose.
### Method Composition
Currently, methods are monolithic classes. Should we support composing methods
(e.g., DFSFT pre-training followed by Self-Forcing distillation) within a
single config? Or is sequential training with checkpoint handoff sufficient?
### New Models and Methods
What models and training methods should we prioritize next?
- **Models:** HunyuanVideo, CogVideoX, other Wan variants?
- **Methods:** Consistency models, progressive distillation, reward-based
fine-tuning?
### Distributed Strategy
Currently supports HSDP (hybrid sharded data parallel) + TP + SP. Are there
scenarios where the current distributed setup is insufficient? Should we add
pipeline parallelism for very large models?
---
## References
- [Self-Forcing paper](https://arxiv.org/abs/2406.05477) — Chen et al., 2024.
- [DMD2 paper](https://arxiv.org/abs/2405.14867) — Yin et al., 2024.
- [Diffusion Forcing paper](https://arxiv.org/abs/2407.01392) — Chen et al.,
2024.
+18 -15
View File
@@ -8,17 +8,9 @@ FastVideo supports the following hardware platforms:
## Quick Installation
### Using pip
### Using uv (recommended)
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
pip install fastvideo
```
### Using uv
Use uv as the default environment manager for faster and more stable installs.
```bash
# Create and activate a new uv environment
@@ -28,21 +20,32 @@ source .venv/bin/activate
uv pip install fastvideo
```
### Using Conda (alternative)
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
pip install fastvideo
```
### From source
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
pip install -e .
# or if you are using uv
uv pip install -e .
# optional: install flash-attn
uv pip install flash-attn --no-build-isolation -v
```
Also optionally install flash-attn:
Alternative with Conda environment:
```bash
pip install flash-attn --no-build-isolation
pip install -e .
pip install flash-attn --no-build-isolation -v
```
## Hardware Requirements
+44 -22
View File
@@ -12,8 +12,20 @@ Instructions to install FastVideo for NVIDIA CUDA GPUs.
## Set up using Python
### Create a new Python environment
#### Conda
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
#### uv
Recommended default: use [uv](https://docs.astral.sh/uv/) for faster and more stable environment setup.
Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, create a new environment using:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools`.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
#### Conda (alternative)
You can also create a Python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html).
##### 1. Install Miniconda (if not already installed)
```bash
@@ -25,34 +37,35 @@ source ~/.bashrc
##### 2. Create and activate a Conda environment for FastVideo
```bash
# (Recommended) Create a new conda environment.
# Create and activate a Conda environment
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
#### uv
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
### Installation
```bash
pip install fastvideo
#### With uv (recommended)
# or if you are using uv
```bash
uv pip install fastvideo
```
Also optionally install flash-attn:
Also optionally install FlashAttention:
```bash
pip install flash-attn --no-build-isolation
uv pip install flash-attn --no-build-isolation -v
```
#### With Conda environment (alternative)
```bash
pip install fastvideo
```
Also optionally install FlashAttention:
```bash
pip install flash-attn --no-build-isolation -v
```
### Installation from Source
@@ -68,18 +81,27 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
pip install -e .
# or if you are using uv
uv pip install -e .
```
Alternative with Conda environment:
```bash
pip install -e .
```
### Optional Dependencies
#### Flash Attention
```bash
pip install flash-attn --no-build-isolation
uv pip install flash-attn --no-build-isolation -v
```
Alternative with Conda environment:
```bash
pip install flash-attn --no-build-isolation -v
```
## Set up using Docker
+27 -19
View File
@@ -11,9 +11,20 @@ Instructions to install FastVideo for Apple Silicon.
### Create a new Python environment
#### Conda
#### uv
Recommended default: use [uv](https://docs.astral.sh/uv/) for faster and more stable environment setup.
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, create a new environment using:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools`.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
#### Conda (alternative)
You can also create a Python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html).
##### 1. Install Miniconda (if not already installed)
@@ -26,21 +37,10 @@ source ~/.zshrc
##### 2. Create and activate a Conda environment for FastVideo
```bash
# (Recommended) Create a new conda environment.
conda create -n fastvideo python=3.12.4 -y
conda activate fastvideo
```
#### uv
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
### Dependencies
```
@@ -49,11 +49,16 @@ brew install ffmpeg
### Installation
#### With uv (recommended)
```bash
uv pip install fastvideo
```
#### With Conda environment (alternative)
```bash
pip install fastvideo
# or if you are using uv
uv pip install fastvideo
```
### Installation from Source
@@ -69,12 +74,15 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
pip install -e .
# or if you are using uv
uv pip install -e .
```
Alternative with Conda environment:
```bash
pip install -e .
```
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
+5 -6
View File
@@ -7,18 +7,18 @@ Get up and running with FastVideo in minutes!
First, install FastVideo:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# If you previously used Conda, use uv instead for a faster, more stable setup
uv venv --python 3.12 --seed
source .venv/bin/activate
# Install FastVideo
pip install fastvideo
uv pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
uv pip install flash-attn --no-build-isolation -v
```
## Basic Usage
@@ -41,7 +41,6 @@ def main():
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
+2 -3
View File
@@ -26,8 +26,7 @@ FastVideo is an inference and post-training framework for diffusion models. It f
FastVideo has the following features:
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sliding Tile Attention](attention/sta/index.md)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- E2E post-training support
- Data preprocessing pipeline for video data
@@ -45,7 +44,7 @@ Use the navigation menu on the left to explore different sections:
- **Inference**: Learn how to use FastVideo for video generation
- **Training**: Data preprocessing and fine-tuning workflows
- **Distillation**: Post-training optimization techniques
- **Sliding Tile Attention**: Advanced attention mechanisms
- **Sliding Tile Attention**: Legacy workflow docs and kernel notes
- **Video Sparse Attention**: Efficient attention for video models
- **Design**: Framework architecture and design principles
- **Developer Guide**: Contributing and development setup
+1 -1
View File
@@ -123,7 +123,7 @@ self.attn = DistributedAttention(
softmax_scale=None,
causal=False,
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
+460
View File
@@ -0,0 +1,460 @@
# Inference Architecture
This section documents the FastVideo inference pipeline: how models are
discovered, configs resolved, components loaded, and stages composed to
generate video. Training-specific code paths (FSDP, gradient checkpointing)
are out of scope.
## Registries
FastVideo uses three registries that work together to resolve a
user-provided `model_path` into a runnable pipeline.
### Model Registry (`fastvideo/models/registry.py`)
Maps HuggingFace architecture class names (e.g. `"WanTransformer3DModel"`)
to FastVideo model classes. Two discovery mechanisms:
1. **Hardcoded dicts** — `_TEXT_TO_VIDEO_DIT_MODELS`, `_VAE_MODELS`,
`_SCHEDULERS`, `_TEXT_ENCODER_MODELS`, `_IMAGE_ENCODER_MODELS`,
`_UPSAMPLERS`, `_AUDIO_MODELS`. Each entry is
`{hf_class_name: (component_name, module_name, class_name)}`.
2. **AST-based discovery** — `_discover_and_register_models()` walks
`fastvideo/models/` and parses each `.py` file's AST looking for an
`EntryClass` variable assignment. Discovered models take priority over
hardcoded entries. For example,
`fastvideo/models/dits/wanvideo.py` exports
`EntryClass = WanTransformer3DModel`.
Both feed into a unified `_FAST_VIDEO_MODELS` dict, which populates the
singleton `ModelRegistry` — an instance of `_ModelRegistry`. Components are
wrapped in `_LazyRegisteredModel` for deferred import.
**Key API:** `ModelRegistry.resolve_model_cls(architectures)` iterates
candidate architecture strings and returns the first matching
`(model_cls, arch)` tuple. Called by `TransformerLoader`, `VAELoader`, etc.
### Config Registry (`fastvideo/registry.py`)
Maps model paths and names to `(PipelineConfig, SamplingParam)` class
pairs. Registration happens at module load via `_register_configs()`, which
calls `register_configs()` for each model family:
```python
register_configs(
sampling_param_cls=WanT2V_1_3B_SamplingParam,
pipeline_config_cls=WanT2V480PConfig,
hf_model_paths=["Wan-AI/Wan2.1-T2V-1.3B-Diffusers"],
model_detectors=[lambda path: "wanpipeline" in path.lower()],
)
```
Each call populates three data structures:
- `_CONFIG_REGISTRY: dict[str, ConfigInfo]` — auto-incrementing ID to
`ConfigInfo(sampling_param_cls, pipeline_config_cls)`.
- `_MODEL_HF_PATH_TO_NAME: dict[str, str]` — HF path to registry ID.
- `_MODEL_NAME_DETECTORS: list[tuple[str, Callable]]` — lambda detectors.
**Resolution priority** (`_get_config_info()`):
1. Exact HF path match in `_MODEL_HF_PATH_TO_NAME`.
2. Partial match on short model name (last path segment, case-insensitive).
3. Detector-based match — runs each detector against the lowercased path
and the `_class_name` from `model_index.json`.
4. `RuntimeError` if no match.
**Top-level resolver:** `get_model_info(model_path, pipeline_type,
workload_type)` combines config resolution with pipeline resolution to
return a `ModelInfo(pipeline_cls, sampling_param_cls, pipeline_config_cls)`.
### Pipeline Registry (`fastvideo/pipelines/pipeline_registry.py`)
Discovers pipeline classes by scanning Python packages under
`fastvideo/pipelines/{basic,preprocess,training}/`.
`import_pipeline_classes()` iterates architecture subdirectories
(e.g. `wan/`, `hunyuan/`), imports each module, and collects those
exporting an `EntryClass` attribute. Supports single class or list.
Returns `{pipeline_type_str: {pipeline_class_name: pipeline_cls}}`.
`_PipelineRegistry.resolve_pipeline_cls(pipeline_name, pipeline_type,
workload_type)` looks up the pipeline class by the `_class_name` field
from `model_index.json`.
## Config Mechanism
### Config Hierarchy
```
PipelineConfig (fastvideo/configs/pipelines/base.py)
├── WanT2V480PConfig (fastvideo/configs/pipelines/wan.py)
│ ├── WanT2V720PConfig
│ └── WanI2V480PConfig
├── HunyuanConfig (fastvideo/configs/pipelines/hunyuan.py)
├── LTX2T2VConfig (fastvideo/configs/pipelines/ltx2.py)
├── CosmosConfig (fastvideo/configs/pipelines/cosmos.py)
└── ... (15+ model families)
```
`PipelineConfig` holds:
- Video generation params: `embedded_cfg_scale`, `flow_shift`,
`disable_autocast`, `is_causal`.
- Nested model configs: `dit_config: DiTConfig`, `vae_config: VAEConfig`,
`text_encoder_configs: tuple[EncoderConfig, ...]`.
- Precision settings: `dit_precision`, `vae_precision`,
`text_encoder_precisions`.
Model-specific subclasses override defaults. For example,
`WanT2V480PConfig` sets `flow_shift=3.0` and uses `WanVideoConfig` as
its DiT config.
### ModelConfig / ArchConfig (`fastvideo/configs/models/base.py`)
`ModelConfig` wraps an `ArchConfig` using `__getattr__` proxy — attribute
access falls through to `arch_config` transparently. `ArchConfig` holds
architecture fields from `config.json` (hidden_size, num_attention_heads,
etc.) and is immutable after initialization. `update_model_arch()` writes
to `ArchConfig`; `update_model_config()` writes to `ModelConfig` fields.
Concrete hierarchy: `DiTConfig` → `DiTArchConfig`, `VAEConfig` →
`VAEArchConfig`, `EncoderConfig` → `EncoderArchConfig`.
### Config Construction
- `PipelineConfig.from_pretrained(model_path)` — resolves config class
via `get_pipeline_config_cls_from_name()`, instantiates with defaults.
- `PipelineConfig.from_kwargs(kwargs)` — resolves class, optionally loads
JSON via `load_from_json()`, then applies CLI overrides via
`update_config_from_dict()`.
- `dump_to_json()` / `load_from_json()` — JSON persistence. Callable
fields and `arch_config` are excluded from dumps.
### SamplingParam (`fastvideo/configs/sample/`)
Generation parameters separate from pipeline config. Each model family
provides defaults:
```python
@dataclass
class WanT2V_1_3B_SamplingParam(SamplingParam):
height: int = 480
width: int = 832
num_frames: int = 81
guidance_scale: float = 3.0
num_inference_steps: int = 50
```
## Component Loading
### ComponentLoader (`fastvideo/models/loader/component_loader.py`)
Abstract base with a `load(model_path, fastvideo_args)` method.
`ComponentLoader.for_module_type(module_type, library)` is a factory
that dispatches to specialized loaders via a `module_loaders` dict:
| Module type | Loader class | Library |
|---|---|---|
| `scheduler` | `SchedulerLoader` | diffusers |
| `transformer`, `transformer_2`, `transformer_3` | `TransformerLoader` | diffusers |
| `vae` | `VAELoader` | diffusers |
| `text_encoder`, `text_encoder_2`, `text_encoder_3` | `TextEncoderLoader` | transformers |
| `tokenizer`, `tokenizer_2`, `tokenizer_3` | `TokenizerLoader` | transformers |
| `image_encoder` | `ImageEncoderLoader` | transformers |
| `image_processor`, `feature_extractor` | `ImageProcessorLoader` | transformers |
| `audio_vae`, `audio_decoder` | `AudioDecoderLoader` | diffusers |
| `vocoder` | `VocoderLoader` | diffusers |
| `upsampler`, `upsampler_2` | `UpsamplerLoader` | diffusers |
`TransformerLoader` reads `config.json` from the component directory,
resolves the class via `ModelRegistry.resolve_model_cls()`, instantiates
the model, and loads safetensors weights. CPU offload and layerwise
offload are applied based on `FastVideoArgs`.
Unknown module types fall back to `GenericComponentLoader`.
### model_index.json
Diffusers-format JSON at the model root. Keys are module names; values
are `[library, class_name]` tuples:
```json
{
"_class_name": "WanPipeline",
"_diffusers_version": "0.24.0",
"transformer": ["diffusers", "WanTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"]
}
```
`ComposedPipelineBase.load_modules()` reads this file via
`_load_config()`, strips metadata keys (`_class_name`,
`_diffusers_version`, `_name_or_path`), detects MoE pipelines via
`boundary_ratio`, then loads each module listed in the pipeline's
`required_config_modules`. Modules not in `required_config_modules` are
skipped. `_extra_config_module_map` allows aliasing (e.g. mapping
`"transformer_2"` to an alternate directory name).
`PipelineComponentLoader.load_module()` orchestrates per-component
loading by calling `ComponentLoader.for_module_type()` then `.load()`.
## Stage Design
### PipelineStage (`fastvideo/pipelines/stages/base.py`)
Abstract base class using the Template Method pattern:
- `__call__(batch, fastvideo_args)` — orchestrates verification, timing,
and error handling. Not overridden by subclasses.
- `forward(batch, fastvideo_args) -> ForwardBatch` — abstract, contains
the stage logic.
- `verify_input()` / `verify_output()` — optional hooks returning
`VerificationResult`. Default: no checks.
When `fastvideo_args.enable_stage_verification` is `True`, `__call__`
runs input verification before `forward()` and output verification after.
When `envs.FASTVIDEO_STAGE_LOGGING` is set, execution time is measured
with `torch.cuda.synchronize()` and logged.
### ForwardBatch (`fastvideo/pipelines/pipeline_batch_info.py`)
Dataclass carrying all pipeline state between stages. Key field groups:
- **Inputs**: `prompt`, `negative_prompt`, `image_path`, `pil_image`,
`video_path`.
- **Embeddings**: `prompt_embeds: list[Tensor]`,
`negative_prompt_embeds`, `prompt_attention_mask`, `image_embeds`.
- **Latents**: `latents`, `image_latent`, `noise_pred`,
`lq_latents`.
- **Dimensions**: `height`, `width`, `num_frames`, `height_latents`,
`width_latents`.
- **Scheduler**: `timesteps`, `num_inference_steps`, `guidance_scale`,
`sigmas`.
- **Task-specific**: `mouse_cond`/`keyboard_cond` (MatrixGame), `pose`
(HYWorld), `camera_states` (GameCraft), `c2ws_plucker_emb`
(LingBotWorld).
- **Output**: `output: Tensor | None`.
- **Logging**: `logging_info: PipelineLoggingInfo`.
`__post_init__` enables CFG when `guidance_scale > 1.0` or LTX2 text
CFG scales differ from 1.0.
### Stage Catalog
Standard stages (typical execution order):
| Stage | File | Purpose |
|---|---|---|
| `InputValidationStage` | `stages/input_validation.py` | Validates input dimensions and types |
| `TextEncodingStage` | `stages/text_encoding.py` | Encodes prompts via text encoders |
| `ImageEncodingStage` | `stages/image_encoding.py` | Encodes input images (I2V pipelines) |
| `ConditioningStage` | `stages/conditioning.py` | Prepares conditioning embeddings |
| `TimestepPreparationStage` | `stages/timestep_preparation.py` | Sets up scheduler timesteps |
| `LatentPreparationStage` | `stages/latent_preparation.py` | Initializes noise latents |
| `DenoisingStage` | `stages/denoising.py` | Main diffusion denoising loop |
| `DecodingStage` | `stages/decoding.py` | Decodes latents to video via VAE |
Specialized variants: `CausalDenoisingStage`, `LTX2DenoisingStage`,
`LongCatDenoisingStage`, `GameCraftDenoisingStage`,
`HYWorldDenoisingStage`, `MatrixGameDenoisingStage`,
`SRDenoisingStage`, `LTX2AudioDecodingStage`, `SD35ConditioningStage`,
`LTX2TextEncodingStage`, `LTX2LatentPreparationStage`.
### Verification System (`fastvideo/pipelines/stages/validators.py`)
`StageValidators` (aliased as `V`) provides static validators:
`not_none`, `positive_int`, `is_tensor`, `tensor_with_dims`,
`positive_int_divisible(divisor)`, etc.
`VerificationResult` collects check results:
```python
result = VerificationResult()
result.add_check("height", batch.height, V.positive_int_divisible(8))
result.add_check("width", batch.width, V.positive_int_divisible(8))
```
`is_valid()` returns whether all checks passed. `get_failure_summary()`
provides detailed error messages. Failed verification raises
`StageVerificationError`.
## Pipeline Architecture
### ComposedPipelineBase (`fastvideo/pipelines/composed_pipeline_base.py`)
Abstract base for all inference pipelines. Lifecycle:
1. **`__init__(model_path, fastvideo_args)`** — initializes distributed
environment via `maybe_init_distributed_environment_and_model_parallel
(tp_size, sp_size)`, then calls `load_modules()` to populate
`self.modules`.
2. **`post_init()`** — calls `initialize_pipeline()` (model-specific
setup), `create_pipeline_stages()` (abstract — subclasses wire stages),
optionally applies `torch.compile` to transformers, and calls
`warmup_sequence_parallel_communication()`.
3. **`forward(batch, fastvideo_args)`** — iterates `self.stages` calling
each stage in order. Decorated with `@torch.no_grad()`.
Key class attributes:
- `_required_config_modules: list[str]` — module names to load from
`model_index.json`.
- `_extra_config_module_map: dict[str, str]` — aliases for module dirs.
- `is_video_pipeline: bool` — whether this produces video output.
Key methods:
- `add_stage(name, stage)` — appends to `_stages` list and
`_stage_name_mapping` dict, also sets attribute on `self`.
- `get_module(name, default)` — retrieves a loaded module.
- `from_pretrained(model_path, **kwargs)` — class method constructing
`FastVideoArgs` and calling `cls(...)` then `post_init()`.
### LoRAPipeline (`fastvideo/pipelines/lora_pipeline.py`)
Extends `ComposedPipelineBase` with LoRA adapter support. Sits in the
MRO between the concrete pipeline and `ComposedPipelineBase`:
```python
class WanPipeline(LoRAPipeline, ComposedPipelineBase):
...
```
Key functionality:
- `convert_to_lora_layers()` — scans transformer blocks, replaces target
linear layers (default: q/k/v/o projections) with LoRA equivalents via
`get_lora_layer()`.
- `set_lora_adapter(path)` — loads safetensors containing `lora_A`,
`lora_B`, `lora_alpha` and maps weights to internal layers.
- `merge_lora_weights()` / `unmerge_lora_weights()` — activates or
deactivates LoRA in the forward pass.
- `LoRAModelLayers` — groups LoRA layers by transformer block for
efficient layerwise offload.
### Distributed Inference
`maybe_init_distributed_environment_and_model_parallel(tp_size, sp_size)`
in `fastvideo/distributed/` initializes `torch.distributed` and creates
tensor-parallel (TP) and sequence-parallel (SP) process groups.
Key APIs: `get_tp_rank()`, `get_tp_world_size()`, `get_sp_rank()`,
`get_sp_world_size()`, `get_world_rank()`, `get_world_size()`.
`warmup_sequence_parallel_communication()` pre-warms NCCL communicators
to avoid slow first forward passes.
Usage: `torchrun --nproc-per-node=N -m fastvideo.entrypoints.cli.main
generate --model-path ... --tp-size N --sp-size M`.
### torch.compile Integration
When `fastvideo_args.enable_torch_compile` is `True`,
`_maybe_compile_pipeline_module()` checks for a `_compile_conditions`
attribute on the module. If present, only matching submodules are
compiled. Otherwise, the entire module is compiled. FSDP-wrapped
modules are skipped.
### Entry Points
**Python API** (`fastvideo/entrypoints/video_generator.py`):
```python
generator = VideoGenerator.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-14B-Diffusers",
num_gpus=1, tp_size=1, sp_size=1,
)
result = generator.generate_video(
prompt="A cat dancing",
height=720, width=1280, num_frames=81,
)
```
**CLI** (`fastvideo/entrypoints/cli/`):
```bash
fastvideo generate \
--model-path "Wan-AI/Wan2.1-T2V-14B-Diffusers" \
--prompt "A cat dancing" \
--num-gpus 1
```
**FastVideoArgs** (`fastvideo/fastvideo_args.py`): Central args dataclass.
Key fields: `model_path`, `mode` (`ExecutionMode`), `workload_type`
(`WorkloadType`), `pipeline_config` (`PipelineConfig`), `num_gpus`,
`tp_size`, `sp_size`, `lora_path`, `dit_cpu_offload`,
`dit_layerwise_offload`, `enable_torch_compile`,
`enable_stage_verification`.
Constructed via `FastVideoArgs.from_kwargs(**kwargs)` which resolves the
`PipelineConfig` from the registry, applies JSON config if provided, and
merges CLI overrides.
## End-to-End Inference Flow
```
User: VideoGenerator.from_pretrained(model_path, **kwargs)
│
├─ FastVideoArgs.from_kwargs() → PipelineConfig resolved via registry
├─ get_model_info() → ModelInfo(pipeline_cls, sampling_param_cls, ...)
│ ├─ model_index.json read → _class_name extracted
│ ├─ pipeline_registry resolves pipeline_cls from _class_name
│ └─ config_registry resolves config classes from model_path
│
├─ pipeline_cls.__init__(model_path, fastvideo_args)
│ ├─ maybe_init_distributed(tp_size, sp_size)
│ └─ load_modules() → reads model_index.json, loads each component
│ ├─ ComponentLoader.for_module_type() → specialized loader
│ └─ loader.load() → model class resolved, weights loaded
│
└─ pipeline.post_init()
├─ initialize_pipeline() → model-specific setup
├─ create_pipeline_stages() → stages wired with modules
├─ torch.compile (if enabled)
└─ warmup_sequence_parallel_communication()
User: generator.generate_video(prompt, ...)
│
├─ ForwardBatch constructed from SamplingParam + user args
└─ pipeline.forward(batch, fastvideo_args)
├─ InputValidationStage → validates dims
├─ TextEncodingStage → prompt → embeddings
├─ ConditioningStage → prepares conditioning
├─ TimestepPreparationStage → scheduler timesteps
├─ LatentPreparationStage → random noise
├─ DenoisingStage → iterative denoising loop
└─ DecodingStage → latents → video frames
```
## Adding a New Model Family — Checklist
1. **Pipeline config** — Create a `PipelineConfig` subclass in
`fastvideo/configs/pipelines/<model>.py`. Set DiT/VAE/encoder configs,
flow_shift, precision defaults.
2. **Sampling param** — Create a `SamplingParam` subclass in
`fastvideo/configs/sample/<model>.py`. Set default height, width,
num_frames, guidance_scale, num_inference_steps.
3. **Register configs** — In `fastvideo/registry.py`, add a
`register_configs()` call inside `_register_configs()` with
`hf_model_paths` and/or `model_detectors`.
4. **Pipeline class** — Create a subclass of `ComposedPipelineBase` (or
`LoRAPipeline` + `ComposedPipelineBase`) in
`fastvideo/pipelines/basic/<model>/<model>_pipeline.py`.
- Set `_required_config_modules` listing needed components.
- Implement `create_pipeline_stages()` wiring stages via `add_stage()`.
- Optionally override `initialize_pipeline()` for custom setup.
- Export `EntryClass = YourPipeline` at module level.
5. **Model classes** (if custom) — Add DiT/VAE implementations in
`fastvideo/models/dits/` or `fastvideo/models/vaes/` with
`EntryClass = YourModel`. The AST discovery will register them
automatically.
6. **Custom stages** (if needed) — Subclass `PipelineStage` in
`fastvideo/pipelines/stages/`, implement `forward()`, optionally
implement `verify_input()`/`verify_output()`.
7. **Verify** — Run `fastvideo generate --model-path <path> --prompt
"test" --num-inference-steps 2` to confirm the pipeline loads and
generates output.
+4 -3
View File
@@ -61,12 +61,11 @@ def main():
prompt,
sampling_param=sampling_param,
output_path="my_videos/", # Controls where videos are saved
return_frames=True, # Also return frames from this call (defaults to False)
save_video=True
)
# If return_frames=True, video contains the generated frames as a NumPy array
print(f"Generated {len(video)} frames")
# If return_frames=True, frames are available in video["frames"]
print(f"Generated {len(video['frames'])} frames")
if __name__ == '__main__':
main()
@@ -76,6 +75,8 @@ if __name__ == '__main__':
The CLI supports `--config` with JSON or YAML. Command-line arguments override
config file values.
By default, `fastvideo generate` uses `return_frames=false` unless you set
`--return-frames` (or `return_frames: true` in config).
```bash
fastvideo generate --config config.yaml
+5 -6
View File
@@ -11,15 +11,15 @@ This page contains step-by-step instructions to get you quickly started with vid
## Installation
We recommend using an environment manager such as `Conda` to create a clean environment:
If you previously used Conda, we recommend using [uv](https://docs.astral.sh/uv/) instead for a faster and more stable environment setup:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Create and activate a new uv environment
uv venv --python 3.12 --seed
source .venv/bin/activate
# Install FastVideo
pip install fastvideo
uv pip install fastvideo
```
For advanced installation options, see the [Installation Guide](../getting_started/installation.md).
@@ -44,7 +44,6 @@ def main():
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
+29 -62
View File
@@ -8,26 +8,24 @@ This page describes the various options for speeding up generation times in Fast
- Optimized Attention Backends
- [Flash Attention](#flash-attention)
- [Sliding Tile Attention](#sliding-tile-attention)
- [Sliding Tile Attention (Archived)](#sliding-tile-attention-archived)
- [Sage Attention](#sage-attention)
- [Sage Attention 3](#sage-attention-3)
- Caching Techniques
- [TeaCache](#teacache)
## Attention Backends
### Available Backends
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
- Video MoBA Attention: `FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN`
- Sparse Linear Attention: `FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
- SageSLA Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_SLA_ATTN`
- Sliding Tile Attention (archived branch only):
`FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
### Configuring Backends
@@ -38,7 +36,7 @@ There are two ways to configure the attention backend in FastVideo.
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
```python
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
```
#### 2. In CLI
@@ -69,12 +67,20 @@ pip install ninja
python setup.py install
```
### Sliding Tile Attention
### Sliding Tile Attention (Archived)
**`SLIDING_TILE_ATTN`**
Sliding Tile Attention is provided by `fastvideo-kernel`.
See [STA docs](../attention/sta/index.md) for installation details.
The full STA integration in `fastvideo/` is archived from `main` and preserved
at:
- https://github.com/hao-ai-lab/FastVideo/tree/sta_do_not_delete
We keep STA off `main` because we believe VSA is strictly better than STA for
the actively maintained FastVideo path.
Kernel code in `fastvideo-kernel` is still retained. For mask search and STA
inference workflow, see [STA docs](../attention/sta/index.md).
### Video Sparse Attention
@@ -117,63 +123,24 @@ These backends are model-specific and require the corresponding kernels and
dependencies. Use the support matrix and model examples to confirm compatibility
before enabling them.
## Teacache
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
### What is TeaCache?
See the official [TeaCache](https://github.com/ali-vilab/TeaCache) repo and their paper for more details.
### How to Enable TeaCache
Enabling TeaCache is straightforward - simply add the `enable_teacache=True` parameter to your `generate_video()` call:
```python
# ... previous code
generator.generate_video(
prompt="Your prompt here",
sampling_param=params,
enable_teacache=True
)
# more code ...
```
### Complete Example
At the bottom is a complete example of using TeaCache for faster video generation. You can run it using the following command:
```bash
python examples/inference/optimizations/teacache_example.py
```
### Advanced Configuration
While TeaCache works well with default settings, you can fine-tune its behavior by adjusting the threshold value:
1. Lower threshold values (e.g., 0.1) will result in more skipped calculations and faster generation with slightly more potential for quality degradation
2. Higher threshold values (e.g., 0.15-0.23) will skip fewer calculations but maintain quality closer to the original
Note that the optimal threshold depends on your specific model and content.
## Benchmarking different optimizations
To benchmark the performance improvement, try generating the same video with and without TeaCache enabled and compare the generation times:
To benchmark backend performance, generate the same prompt with the same seed and compare end-to-end generation times:
```python
# Without TeaCache
start_time = time.perf_counter()
generator.generate_video(prompt="Your prompt", enable_teacache=False)
standard_time = time.perf_counter() - start_time
import os
import time
# With TeaCache
start_time = time.perf_counter()
generator.generate_video(prompt="Your prompt", enable_teacache=True)
teacache_time = time.perf_counter() - start_time
print(f"Standard generation: {standard_time:.2f} seconds")
print(f"TeaCache generation: {teacache_time:.2f} seconds")
print(f"Speedup: {standard_time/teacache_time:.2f}x")
for backend in ["TORCH_SDPA", "FLASH_ATTN", "SAGE_ATTN"]:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = backend
generator = VideoGenerator.from_pretrained("your-model-id")
start_time = time.perf_counter()
generator.generate_video(
prompt="Your prompt",
seed=1024,
)
elapsed = time.perf_counter() - start_time
print(f"{backend}: {elapsed:.2f}s")
```
Note: If you want to benchmark different attention backends, you'll need to reinstantiate `VideoGenerator`.
Note: reinstantiate `VideoGenerator` after changing `FASTVIDEO_ATTENTION_BACKEND`.
+14 -2
View File
@@ -6,6 +6,13 @@ For the canonical, code-level list of model IDs recognized by
`VideoGenerator.from_pretrained(...)`, see the registrations in
`fastvideo/registry.py` (`register_configs(...)` entries).
!!! note
The full STA integration in `fastvideo/` is archived from `main` and kept
in `sta_do_not_delete`:
https://github.com/hao-ai-lab/FastVideo/tree/sta_do_not_delete
We do this because we believe VSA is strictly better than STA for the
actively maintained `main` inference path.
The symbols used have the following meanings:
- ✅ = Full compatibility
@@ -46,7 +53,7 @@ pipeline initialization and sampling.
}
</style>
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA | BSA |
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn (Legacy Branch) | Sage Attn | VSA | BSA |
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
@@ -69,6 +76,9 @@ pipeline initialization and sampling.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
`Sliding Tile Attn (Legacy Branch)` entries refer to the archived
`sta_do_not_delete` branch workflow, not active `main` inference wiring.
## Canonical Supported IDs
The authoritative source for model-ID recognition is
@@ -78,7 +88,9 @@ resolve default pipeline and sampling configuration for it.
## Special requirements
### Sliding Tile Attention
- Currently only Hopper GPUs (H100s) are supported.
- Full STA pipeline usage is on the archived branch:
https://github.com/hao-ai-lab/FastVideo/tree/sta_do_not_delete
- STA currently requires Hopper GPUs (H100s).
### TurboWan2.1 (TurboDiffusion)
- Uses TurboDiffusionPipeline with RCM scheduler for 1-4 step generation
+595
View File
@@ -0,0 +1,595 @@
# Training Infrastructure
FastVideo's training infrastructure (`fastvideo/train/`) is a YAML-driven
framework for training and distilling video diffusion models. A single config
file controls everything — models, algorithms, distributed strategy,
checkpointing, and validation — with no code changes needed to mix and match.
!!! note "Relationship to legacy training"
This system replaces the older script-based training in `fastvideo/training/`.
The legacy scripts still work for basic fine-tuning, but new development
should use the config-driven system documented here.
---
## Quick Start
### Launch with the helper script
```bash
bash examples/train/run.sh examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml
```
The script auto-detects available GPUs and sets up `torchrun`. Override with
environment variables:
```bash
NUM_GPUS=4 NNODES=2 NODE_RANK=0 \
MASTER_ADDR=10.0.0.1 MASTER_PORT=29501 \
bash examples/train/run.sh my_config.yaml
```
### Launch directly with torchrun
```bash
torchrun --nproc_per_node=8 \
fastvideo/train/entrypoint/train.py \
--config examples/train/distill_wan2.1_t2v_1.3B_dmd2.yaml
```
### CLI flags
| Flag | Description |
|------|-------------|
| `--config` | Path to YAML config file (required) |
| `--resume-from-checkpoint` | Path to a DCP checkpoint directory to resume from |
| `--override-output-dir` | Override `training.checkpoint.output_dir` |
| `--dry-run` | Validate config and exit without training |
---
## Config Format
Every run is defined by a single YAML file with five top-level sections.
See `examples/train/example.yaml` for a fully-commented reference.
### `models` — Role-based model instances
Each entry defines a model role. The `_target_` field specifies the Python class
to instantiate:
```yaml
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
```
Common model parameters:
| Parameter | Default | Description |
|-----------|---------|-------------|
| `_target_` | *(required)* | Python class path for the model |
| `init_from` | *(required)* | HuggingFace repo ID or local checkpoint path |
| `trainable` | `true` | Whether the model's parameters require gradients |
| `disable_custom_init_weights` | `false` | Skip custom weight initialization (use for teacher/critic) |
| `flow_shift` | `3.0` | Timestep shifting factor |
| `enable_gradient_checkpointing_type` | `null` | Gradient checkpointing (`"full"` or `null`) |
Which roles are needed depends on the training method:
| Method | Required roles |
|--------|---------------|
| Fine-tune (SFT) | `student` |
| Diffusion-Forcing SFT | `student` |
| DMD2 | `student`, `teacher`, `critic` |
| Self-Forcing | `student` (causal), `teacher`, `critic` |
### `method` — Training algorithm
Selects and configures the training algorithm:
```yaml
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
rollout_mode: simulate
dmd_denoising_steps: [1000, 750, 500, 250]
generator_update_interval: 5
```
To switch algorithms, change `_target_` and adjust the method-specific keys.
See [Training Methods](#training-methods) for details on each algorithm.
### `training` — Typed infrastructure config
This section maps to typed dataclasses with defaults and validation:
```yaml
training:
distributed:
num_gpus: 8
sp_size: 1 # sequence parallelism
tp_size: 1 # tensor parallelism
hsdp_replicate_dim: 1 # HSDP replication dimension
hsdp_shard_dim: 8 # HSDP sharding dimension
data:
data_path: data/my_dataset
train_batch_size: 1
dataloader_num_workers: 4
training_cfg_rate: 0.1 # classifier-free guidance dropout rate
seed: 1000
num_latent_t: 20
num_height: 448
num_width: 832
num_frames: 77
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant # constant, linear, cosine, polynomial
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/my_run
training_state_checkpointing_steps: 1000 # 0 = disabled
checkpoints_total_limit: 3 # 0 = keep all
tracker:
project_name: my_project
run_name: my_run
model:
weighting_scheme: uniform # uniform, logit_normal, mode
precondition_outputs: false
enable_gradient_checkpointing_type: full
vsa:
sparsity: 0.0 # 0.0 = disabled
decay_rate: 0.0
decay_interval_steps: 0
```
### `callbacks` — Pluggable hooks
Callbacks run at specific points in the training loop (before/after optimizer
steps, at validation time, etc.):
```yaml
callbacks:
grad_clip:
max_grad_norm: 1.0
ema:
_target_: fastvideo.train.callbacks.ema.EMACallback
decay: 0.9999
start_iter: 0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline
dataset_file: path/to/validation.json
every_steps: 100
sampling_steps: [4]
guidance_scale: 5.0
```
See [Callbacks](#callbacks) for details on each callback.
### `pipeline` — Inference pipeline overrides
Optional overrides for the inference pipeline used during validation:
```yaml
pipeline:
flow_shift: 8
```
---
## Training Methods
### Supervised Fine-Tuning (SFT)
Standard flow-matching loss. The simplest method — train the student to predict
noise (or clean x0) from noised data samples.
```yaml
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
attn_kind: dense # "dense" or "vsa"
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `attn_kind` | `"dense"` | Attention mode: `"dense"` (standard) or `"vsa"` (sparse) |
### Diffusion-Forcing SFT (DFSFT)
SFT with **per-chunk inhomogeneous timesteps** — each temporal chunk of the
video gets a different noise level. This is a prerequisite for training causal /
streaming models that must handle mixed-noise inputs.
```yaml
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 3
min_timestep_ratio: 0.0
max_timestep_ratio: 1.0
attn_kind: dense
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `chunk_size` | `3` | Latent frames per temporal chunk |
| `min_timestep_ratio` | `0.0` | Lower bound of timestep sampling range |
| `max_timestep_ratio` | `1.0` | Upper bound of timestep sampling range |
| `attn_kind` | `"dense"` | `"dense"` or `"vsa"` |
### DMD2 (Distribution Matching Distillation)
Distill a many-step teacher into a few-step student. The student learns to match
the teacher's score function, guided by a trainable critic network.
```yaml
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-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.dmd2.DMD2Method
rollout_mode: simulate
dmd_denoising_steps: [1000, 750, 500, 250]
generator_update_interval: 5
real_score_guidance_scale: 4.5
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
```
| Parameter | Default | Description |
|-----------|---------|-------------|
| `rollout_mode` | *(required)* | `"simulate"` (pure noise) or `"data_latent"` (from data) |
| `dmd_denoising_steps` | *(required)* | Timestep schedule for student rollout |
| `generator_update_interval` | `1` | Update student every N critic steps |
| `real_score_guidance_scale` | `1.0` | CFG scale for teacher predictions |
| `fake_score_learning_rate` | *(required)* | Critic optimizer learning rate |
| `fake_score_betas` | *(required)* | Critic optimizer Adam betas |
| `fake_score_lr_scheduler` | *(required)* | Critic LR scheduler type |
### Self-Forcing (Causal DMD)
Extends DMD2 for **streaming / causal video generation**. The student processes
video in temporal chunks, feeding its own denoised outputs as context for future
chunks — simulating autoregressive rollout during training.
Requires a causal model class (e.g., `WanCausalModel`) for the student:
```yaml
models:
student:
_target_: fastvideo.train.models.wan.wan_causal.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
dmd_denoising_steps: [1000, 750, 500, 250]
student_sample_type: sde
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
```
Self-Forcing inherits all DMD2 parameters, plus:
| Parameter | Default | Description |
|-----------|---------|-------------|
| `student_sample_type` | `"sde"` | `"sde"` or `"ode"` for intermediate steps |
| `same_step_across_blocks` | `false` | Use same exit timestep for all blocks |
| `last_step_only` | `false` | Always exit at the final denoising step |
| `context_noise` | `0.0` | Noise added to context frames (0 = clean) |
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
| `start_gradient_frame` | `0` | Frame index where gradients begin |
---
## Callbacks
Callbacks are pluggable hooks that run at specific points in the training loop.
Configure them under the `callbacks` section.
### GradNormClipCallback
Clips gradient norms before the optimizer step. Optionally logs per-module
gradient norms to the tracker.
```yaml
callbacks:
grad_clip:
max_grad_norm: 1.0 # 0.0 = disabled
log_grad_norms: false
```
### EMACallback
Maintains an exponential moving average of the student's weights. The EMA
weights are automatically swapped in during validation.
```yaml
callbacks:
ema:
_target_: fastvideo.train.callbacks.ema.EMACallback
decay: 0.9999
start_iter: 0 # delay EMA updates until this iteration
```
The EMA callback owns its own state and checkpoints independently — EMA weights
are saved and restored automatically on resume.
### ValidationCallback
Runs inference with the trained model at regular intervals, saving generated
videos and logging them to the tracker (W&B).
```yaml
callbacks:
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline
dataset_file: path/to/validation.json
every_steps: 100
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250] # explicit timestep list
guidance_scale: 5.0
rollout_mode: parallel # "parallel" or "streaming"
```
The validation dataset is a JSON file containing a list of prompt strings.
If EMA is enabled, validation automatically uses the EMA weights.
---
## Checkpointing and Resume
### Checkpoint format
Checkpoints use PyTorch Distributed Checkpoint (DCP) format, compatible with
FSDP/HSDP sharding. Each checkpoint saves:
- Model weights (all roles)
- Optimizer states (all roles)
- LR scheduler states
- RNG states (for exact reproducibility)
- EMA shadow weights (if enabled)
- Training step counter
Checkpoints are saved to `<output_dir>/checkpoint-<step>/`.
### Saving checkpoints
```yaml
training:
checkpoint:
output_dir: outputs/my_run
training_state_checkpointing_steps: 1000 # save every N steps (0 = off)
checkpoints_total_limit: 3 # rolling window (0 = keep all)
```
### Resuming training
Use `--resume-from-checkpoint` to resume from a specific checkpoint:
```bash
# Via the helper script
bash examples/train/run.sh my_config.yaml --resume outputs/my_run/checkpoint-2000
# Via torchrun directly
torchrun --nproc_per_node=8 \
fastvideo/train/entrypoint/train.py \
--config my_config.yaml \
--resume-from-checkpoint outputs/my_run/checkpoint-2000
```
Or set it in the YAML:
```yaml
training:
checkpoint:
resume_from_checkpoint: outputs/my_run/checkpoint-2000
```
### Reproducibility
The training entrypoint enables deterministic mode automatically:
- `torch.backends.cudnn.benchmark = False`
- `torch.backends.cudnn.deterministic = True`
- `torch.use_deterministic_algorithms(True)`
A shared CUDA RNG generator is seeded from `training.data.seed` and threaded
through all random operations (noise sampling, timestep sampling, etc.).
Ranks within the same sequence-parallel group share a seed, ensuring identical
noise across SP shards.
---
## Distributed Training
The framework supports HSDP (Hybrid Sharded Data Parallel), Tensor Parallelism
(TP), and Sequence Parallelism (SP):
```yaml
training:
distributed:
num_gpus: 8
sp_size: 1 # sequence parallelism group size
tp_size: 1 # tensor parallelism group size
hsdp_replicate_dim: 1 # number of HSDP replicas
hsdp_shard_dim: 8 # number of HSDP shards
```
**HSDP** shards model parameters across `hsdp_shard_dim` GPUs and replicates
across `hsdp_replicate_dim` groups. The product
`hsdp_replicate_dim * hsdp_shard_dim` should equal `num_gpus`.
**Sequence parallelism** splits the sequence (video frames) across `sp_size`
GPUs within each data-parallel group. Useful for long videos that don't fit on a
single GPU.
---
## VSA (Variable Sparse Attention)
VSA progressively increases attention sparsity during training, reducing compute
while maintaining quality:
```yaml
training:
vsa:
sparsity: 0.9 # target sparsity level
decay_rate: 0.03 # sparsity increment per decay interval
decay_interval_steps: 1 # steps between sparsity increases
```
The effective sparsity at step `t` is
`min(sparsity, decay_rate * (t // decay_interval_steps))`.
---
## Extending the Framework
### Adding a new model
1. Create a new module under `fastvideo/train/models/` (e.g.,
`fastvideo/train/models/mymodel/mymodel.py`).
2. Subclass `ModelBase` (or `CausalModelBase` for streaming models).
3. Implement the required methods:
- `prepare_batch()` — convert raw dataloader output to `TrainingBatch`
- `add_noise()` — forward-process noise addition
- `predict_noise()` — run the transformer forward pass
- `backward()` — backward pass with forward context restoration
4. Reference it in your YAML config:
```yaml
models:
student:
_target_: fastvideo.train.models.mymodel.mymodel.MyModel
init_from: my-org/my-model
trainable: true
```
### Adding a new training method
1. Create a new module under `fastvideo/train/methods/`.
2. Subclass `TrainingMethod`.
3. Implement the required methods:
- `single_train_step()` — one forward pass returning losses, outputs, metrics
- `get_optimizers()` — return optimizer list
- `get_lr_schedulers()` — return scheduler list
4. Reference it in your config:
```yaml
method:
_target_: fastvideo.train.methods.my_method.MyMethod
my_param: 42
```
Method-specific parameters are accessible via `self.method_config` (a plain
dict).
### Adding a new callback
1. Create a new module under `fastvideo/train/callbacks/`.
2. Subclass `Callback`.
3. Override the hooks you need: `on_train_start`, `on_training_step_end`,
`on_before_optimizer_step`, etc.
4. Optionally implement `state_dict()` / `load_state_dict()` for checkpoint
persistence.
5. Add it to your config:
```yaml
callbacks:
my_callback:
_target_: fastvideo.train.callbacks.my_callback.MyCallback
my_param: 42
```
---
## File Structure
```
fastvideo/train/
entrypoint/
train.py # CLI entrypoint (torchrun)
trainer.py # Training loop orchestrator
models/
base.py # ModelBase, CausalModelBase ABCs
wan/
wan.py # Wan 2.1 T2V model
wan_causal.py # Wan causal (streaming) model
methods/
base.py # TrainingMethod ABC
distribution_matching/
dmd2.py # DMD2 distillation
self_forcing.py # Self-Forcing (causal DMD)
fine_tuning/
finetune.py # Supervised fine-tuning
dfsft.py # Diffusion-forcing SFT
callbacks/
callback.py # Callback ABC and CallbackDict
grad_clip.py # Gradient clipping + norm logging
ema.py # EMA weight averaging
validation.py # Periodic inference validation
utils/
config.py # YAML parser -> RunConfig
training_config.py # Typed config dataclasses
builder.py # Model/method instantiation
optimizer.py # Optimizer/scheduler construction
checkpoint.py # DCP save/resume
dataloader.py # Dataset/dataloader construction
tracking.py # W&B tracker
```
---
## Related Docs
- [Training Architecture](../design/training_architecture.md) — design
rationale, model/method abstractions, and open questions.
- [Training Overview](overview.md) — data requirements and preprocessing.
- [Data Preprocessing](data_preprocess.md) — how to prepare datasets.
- [Config Reference](../../examples/train/configs/example.yaml) — fully-commented
YAML config with all fields and defaults.
+3 -1
View File
@@ -49,7 +49,9 @@ combinations.
If forcing a backend fails, verify optional dependencies are installed:
- `FLASH_ATTN`: `flash-attn`
- `SLIDING_TILE_ATTN` and `VIDEO_SPARSE_ATTN`: `fastvideo-kernel`
- `VIDEO_SPARSE_ATTN`: `fastvideo-kernel`
- `SLIDING_TILE_ATTN`: STA legacy workflow in
`sta_do_not_delete` + `fastvideo-kernel`
- `SAGE_ATTN` / `SAGE_ATTN_THREE`: SageAttention packages
As a fallback, use:
+1 -1
View File
@@ -11,7 +11,7 @@ def main():
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
num_gpus=2,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
@@ -168,7 +168,7 @@ def prepare_sampling_params(video_request: VideoGenerationRequest, default_param
params.height = video_request.height
params.width = video_request.width
params.save_video = False
params.return_frames = False
params.return_frames = True
return params
@@ -229,7 +229,7 @@ class BaseModelDeployment:
sampling_param=params,
image_path=image_path,
save_video=False,
return_frames=False,
return_frames=True,
)
inference_time = time.time() - inference_start_time
@@ -3,7 +3,3 @@
```bash
python examples/inference/optimizations/attention_example.py
```
```bash
python examples/inference/optimizations/teacache_example.py
```
@@ -1,47 +0,0 @@
import time
from fastvideo import VideoGenerator, SamplingParam
def main():
start_time = time.perf_counter()
gen = VideoGenerator.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
gen_start_time = time.perf_counter()
params = SamplingParam.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
)
# this controls the threshold for the tea cache
params.teacache_params.teacache_thresh = 0.08
gen.generate_video(
prompt=
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
sampling_param=params,
height=480,
width=832,
num_frames=61, # 85 ,77
num_inference_steps=50,
enable_teacache=True,
seed=1024,
output_path="example_outputs/")
generation_time = time.perf_counter() - gen_start_time
print(f"Video generation time: {generation_time:.2f} seconds")
total_time = time.perf_counter() - start_time
print(f"Total execution time: {total_time:.2f} seconds")
if __name__ == "__main__":
main()
+41 -2
View File
@@ -1,5 +1,44 @@
# STA Mask Search Examples
# STA Mask Search (Archived Workflow)
The full STA integration is preserved in:
- https://github.com/hao-ai-lab/FastVideo/tree/sta_do_not_delete
Switch to the branch before running mask search:
```bash
git fetch origin
git checkout sta_do_not_delete
```
Run mask search + tuning:
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
bash examples/inference/sta_mask_search/inference_wan_sta.sh
```
```
This script runs:
- `STA_searching` and writes to `inference_results/sta/mask_search_full`
- `STA_tuning` and writes to `inference_results/sta/mask_search_sparse`
Example STA inference (same archived branch):
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
fastvideo generate \
--model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \
--num-gpus 2 \
--tp-size 2 \
--sp-size 2 \
--height 768 \
--width 1280 \
--num-frames 69 \
--num-inference-steps 50 \
--prompt "A cinematic wildlife shot of a lion walking in golden grasslands."
```
@@ -1,39 +0,0 @@
#!/bin/bash
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=1
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
output_path="inference_results/sta/mask_search_full"
STA_mode="STA_searching"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode &
sleep 1
done
wait
echo "STA searching completed"
output_path="inference_results/sta/mask_search_sparse"
STA_mode="STA_tuning"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode \
--skip_time_steps $skip_time_steps &
sleep 1
done
wait
echo "STA tuning completed"
echo "All jobs completed"
@@ -1,63 +0,0 @@
import os
import argparse
from fastvideo import VideoGenerator, SamplingParam
def main(args):
os.makedirs(args.output_path, exist_ok=True)
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
num_gpus=args.num_gpus, # Adjust based on your hardware
STA_mode=args.STA_mode,
skip_time_steps=args.skip_time_steps
)
# Prompts for your video
prompt = args.prompt
prompt_path = args.prompt_path
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
if prompt_path is not None:
with open(prompt_path, "r") as f:
prompts = f.readlines()
else:
prompts = [prompt]
params = SamplingParam(
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
fps=args.fps,
guidance_scale=args.guidance_scale,
seed=args.seed,
return_frames=True, # Also return frames from this call (defaults to False)
output_path=args.output_path, # Controls where videos are saved
save_video=True,
negative_prompt=negative_prompt
)
# Generate the video
for prompt in prompts:
video = generator.generate_video(
prompt,
sampling_param=params,
)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--prompt", type=str, default="A man is dancing.")
parser.add_argument("--prompt_path", type=str, default=None)
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1280)
parser.add_argument("--num_frames", type=int, default=69)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--guidance_scale", type=float, default=5.0)
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--output_path", type=str, default="my_videos/")
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--STA_mode", type=str, default="STA_searching")
parser.add_argument("--skip_time_steps", type=int, default=12)
args = parser.parse_args()
main(args)
@@ -0,0 +1,70 @@
# DFSFT (Diffusion-Forcing SFT): Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model
# - Training: inhomogeneous timesteps per chunk (diffusion forcing)
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 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_77x448x832_600k
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: 2.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/wan2.1_causal_dfsft
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_dfsft
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_64.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 8
@@ -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
@@ -0,0 +1,94 @@
# DMD2 distillation: Wan 2.1 T2V 1.3B (teacher 50-step -> student 4-step).
#
# - Teacher: frozen pretrained Wan 2.1 T2V 1.3B
# - Student: trainable, initialized from the same pretrained weights
# - Critic: trainable, initialized from the same pretrained weights
# - Validation: 4-step SDE sampling
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-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.dmd2.DMD2Method
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.5
dmd_denoising_steps: [1000, 750, 500, 250]
# Critic optimizer (required — no fallback to training.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: data/Wan-Syn_77x448x832_600k
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: 2.0e-6
betas: [0.0, 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/wan2.1_dmd2_4steps
training_state_checkpointing_steps: 20
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan
run_name: wan2.1_dmd2_4steps
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_dmd_pipeline.WanDMDPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
guidance_scale: 6.0
ema:
_target_: fastvideo.train.callbacks.ema.EMACallback
decay: 0.98
start_iter: 0 # delay EMA updates until this iteration
pipeline:
flow_shift: 8
+194
View File
@@ -0,0 +1,194 @@
# ==============================================================================
# Full configuration reference for fastvideo.train
#
# Legend:
# [TYPED] — parsed into a typed dataclass; fields are validated with
# defaults. Unknown keys are silently ignored.
# [FREE] — free-form dict passed as-is to the target class / method.
# Keys depend on the _target_ class constructor / method_config.
# [RESOLVED] — parsed by PipelineConfig.from_kwargs(); auto-populated from
# the model's config files. Only scalar overrides are useful.
# ==============================================================================
# ------------------------------------------------------------------------------
# models: [FREE]
#
# Each role is instantiated via _target_(*, training_config=..., **kwargs).
# Keys here are constructor kwargs of the _target_ class (e.g. WanModel).
# You can define any role name (student, teacher, critic, etc.).
# ------------------------------------------------------------------------------
models:
student:
_target_: fastvideo.train.models.wan.WanModel # required
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers # required: HF repo or local path
trainable: true # default: true
disable_custom_init_weights: false # default: false
flow_shift: 3.0 # default: 3.0
enable_gradient_checkpointing_type: null # default: null (falls back to training.model)
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-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: [FREE]
#
# Instantiated via _target_(*, cfg=RunConfig, role_models=...).
# All keys besides _target_ are available in self.method_config (a plain dict).
# Keys depend entirely on the method class.
# ------------------------------------------------------------------------------
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method # required
# --- DMD2-specific keys (read from self.method_config) ---
rollout_mode: simulate # required: "simulate" or "data_latent"
generator_update_interval: 5 # default: 1
dmd_denoising_steps: [1000, 750, 500, 250] # SDE timestep schedule
# Critic optimizer (all required — no fallback)
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
# CFG conditioning policy (optional)
# cfg_uncond:
# on_missing: error # "error" or "ignore"
# text: keep # "keep", "zero", "drop", "negative_prompt"
# image: keep # "keep", "zero", "drop"
# action: keep # "keep", "zero", "drop"
# --- FineTuneMethod keys (if using finetune instead) ---
# _target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
# ------------------------------------------------------------------------------
# training: [TYPED] -> TrainingConfig
#
# Every field below has a typed default. Unknown keys are ignored.
# ------------------------------------------------------------------------------
training:
# --- training.distributed [TYPED] -> DistributedConfig ---
distributed:
num_gpus: 8 # default: 1
tp_size: 1 # default: 1
sp_size: 1 # default: 1 (defaults to num_gpus in loader)
hsdp_replicate_dim: 1 # default: 1
hsdp_shard_dim: 8 # default: -1 (defaults to num_gpus in loader)
pin_cpu_memory: false # default: false
# --- training.data [TYPED] -> DataConfig ---
data:
data_path: data/my_dataset # default: ""
train_batch_size: 1 # default: 1
dataloader_num_workers: 4 # default: 0
training_cfg_rate: 0.1 # default: 0.0
seed: 1000 # default: 0
num_height: 448 # default: 0
num_width: 832 # default: 0
num_latent_t: 20 # default: 0
num_frames: 77 # default: 0
# --- training.optimizer [TYPED] -> OptimizerConfig ---
# Note: only for the student optimizer. Critic optimizer is in method config.
optimizer:
learning_rate: 2.0e-6 # default: 0.0
betas: [0.9, 0.999] # default: [0.9, 0.999]
weight_decay: 0.01 # default: 0.0
lr_scheduler: constant # default: "constant"
lr_warmup_steps: 0 # default: 0
lr_num_cycles: 0 # default: 0
lr_power: 0.0 # default: 0.0
min_lr_ratio: 0.5 # default: 0.5
# --- training.loop [TYPED] -> TrainingLoopConfig ---
loop:
max_train_steps: 10000 # default: 0
gradient_accumulation_steps: 1 # default: 1
# --- training.checkpoint [TYPED] -> CheckpointConfig ---
checkpoint:
output_dir: outputs/my_run # default: ""
resume_from_checkpoint: "" # default: "" (or use --resume-from-checkpoint CLI)
training_state_checkpointing_steps: 1000 # default: 0 (disabled)
checkpoints_total_limit: 3 # default: 0 (keep all)
# --- training.tracker [TYPED] -> TrackerConfig ---
tracker:
trackers: [] # default: [] (auto-adds "wandb" if project_name is set)
project_name: my_project # default: "fastvideo"
run_name: my_run # default: ""
# --- training.vsa [TYPED] -> VSAConfig ---
vsa:
sparsity: 0.0 # default: 0.0 (0.0 = disabled)
# --- training.model [TYPED] -> ModelTrainingConfig ---
model:
weighting_scheme: uniform # default: "uniform"
logit_mean: 0.0 # default: 0.0
logit_std: 1.0 # default: 1.0
mode_scale: 1.0 # default: 1.0
precondition_outputs: false # default: false
moba_config: {} # default: {}
enable_gradient_checkpointing_type: full # default: null ("full" or null)
# --- training top-level [TYPED] ---
dit_precision: fp32 # default: "fp32" (master weight precision)
# model_path: ... # default: "" (auto-derived from models.student.init_from)
# ------------------------------------------------------------------------------
# callbacks: [FREE]
#
# Each callback is instantiated via _target_(*, **kwargs).
# The callback name (e.g. "grad_clip") is arbitrary — only _target_ matters.
# training_config is injected automatically (not from YAML).
# ------------------------------------------------------------------------------
callbacks:
# --- GradNormClipCallback ---
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback # optional if using default registry
max_grad_norm: 1.0 # default: 0.0 (0.0 = disabled)
log_grad_norms: false # default: false
# --- EMACallback ---
# ema:
# _target_: fastvideo.train.callbacks.ema.EMACallback
# decay: 0.9999 # default: 0.9999
# start_iter: 0 # default: 0
# --- ValidationCallback ---
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback # optional if using default registry
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline # required
dataset_file: path/to/validation.json # required
every_steps: 100 # default: 100
sampling_steps: [4] # default: [40]
guidance_scale: 5.0 # default: null (uses model default)
num_frames: null # default: null (derived from training.data)
output_dir: null # default: null (falls back to training.checkpoint.output_dir)
sampling_timesteps: null # default: null (explicit timestep list for streaming)
# ------------------------------------------------------------------------------
# pipeline: [RESOLVED] -> PipelineConfig
#
# Parsed by PipelineConfig.from_kwargs(). Most fields are auto-populated from
# the model's config files (vae_config, dit_config, text_encoder_configs, etc.).
# Only scalar overrides are typically needed here.
# ------------------------------------------------------------------------------
pipeline:
flow_shift: 3 # default: null (model-specific)
# flow_shift_sr: null # default: null (super-resolution shift)
# embedded_cfg_scale: 6.0 # default: 6.0
# is_causal: false # default: false
# vae_tiling: true # default: true
# vae_sp: true # default: true
# disable_autocast: false # default: false
@@ -0,0 +1,74 @@
# V3 config: Wan 2.1 T2V 1.3B finetune with VSA (phase 3.4, 0.9 sparsity).
#
# 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.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/Wan-Syn_77x448x832_600k
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/phase3.4_wan2.1_finetune_vsa_0.9_v3
training_state_checkpointing_steps: 100
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: distillation_wan_r
run_name: phase3.4_wan_finetune_vsa_0.9_v3
model:
enable_gradient_checkpointing_type: full
vsa:
sparsity: 0.9
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: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [50]
guidance_scale: 5.0
pipeline:
flow_shift: 3
@@ -0,0 +1,98 @@
# Self-Forcing distillation: Wan 2.1 T2V 1.3B Causal
#
# - Teacher: frozen pretrained Wan 2.1 T2V 1.3B
# - Student: trainable causal Wan model
# - Critic: trainable, initialized from pretrained weights
# - Training: streaming rollout with SDE sampling
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/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_77x448x832_600k
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: 1e-5
betas: [0.0, 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/wan2.1_causal_self_forcing
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: self-forcing-wan
run_name: wan2.1_causal_self_forcing
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/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 69
guidance_scale: 6.0
pipeline:
flow_shift: 5
@@ -0,0 +1,105 @@
# Self-Forcing distillation: Wan 2.1 T2V 1.3B Causal with VSA (0.9 sparsity)
#
# - Teacher: frozen pretrained Wan 2.1 T2V 14B
# - Student: trainable causal Wan model with VSA sparse attention
# - Critic: trainable, initialized from pretrained weights
# - Training: streaming rollout with SDE sampling + VSA 0.9 sparsity
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/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
attn_kind: vsa
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: 1
training_cfg_rate: 0.0
seed: 1000
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: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_self_forcing_vsa_0.9
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: self-forcing-wan
run_name: wan2.1_causal_self_forcing_vsa_0.9
model:
enable_gradient_checkpointing_type: full
vsa:
sparsity: 0.9
decay_rate: 0.03
decay_interval_steps: 1
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 69
guidance_scale: 6.0
pipeline:
flow_shift: 5
@@ -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
+68
View File
@@ -0,0 +1,68 @@
#!/usr/bin/env bash
# Launch training from a YAML config.
#
# Usage:
# bash examples/train/run.sh <config.yaml> [--dotted.key value ...]
#
# Examples:
# 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/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).
set -euo pipefail
CONFIG="${1:?Usage: $0 <config.yaml> [extra flags...]}"
shift
# ── GPU / node settings ──────────────────────────────────────────
NUM_GPUS="${NUM_GPUS:-$(nvidia-smi -L 2>/dev/null | wc -l)}"
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:-7ff8b6e8356924f7a6dd51a0342dd1a422ea9352}"
export WANDB_MODE="${WANDB_MODE:-online}"
# ── Log file ─────────────────────────────────────────────────────
CONFIG_NAME="$(basename "${CONFIG}" .yaml)"
TIMESTAMP="$(date +%Y%m%d_%H%M%S)"
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}"
echo "Num Nodes: ${NNODES}"
echo "Node Rank: ${NODE_RANK}"
echo "Master: ${MASTER_ADDR}:${MASTER_PORT}"
echo "Extra args: $*"
echo "Log file: ${LOG_FILE}"
echo "=============================="
python -m torch.distributed.run \
--nnodes "${NNODES}" \
--node_rank "${NODE_RANK}" \
--nproc_per_node "${NUM_GPUS}" \
--master_addr "${MASTER_ADDR}" \
--master_port "${MASTER_PORT}" \
fastvideo/train/entrypoint/train.py \
--config "${CONFIG}" \
"$@" \
2>&1 | tee "${LOG_FILE}"
+80
View File
@@ -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[@]}"
+114
View File
@@ -0,0 +1,114 @@
#!/usr/bin/env bash
# Submit a multi-node Slurm training job.
#
# Usage:
# bash examples/train/run_slurm.sh <config.yaml> <num_nodes> [--dotted.key value ...]
#
# Examples:
# bash examples/train/run_slurm.sh examples/train/configs/example.yaml 2
# 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
#
# Environment variables (override defaults):
# PARTITION Slurm partition (default: main)
# NUM_GPUS GPUs per node (default: 8)
# CPUS_PER_TASK CPUs per task (default: 128)
# MEM Memory per node (default: 1440G)
# JOB_NAME Slurm job name (default: derived from config)
# OUTPUT_DIR Directory for slurm logs (default: slurm_logs)
# MASTER_PORT Rendezvous port (default: 29500)
# EXCLUDE Nodes to exclude (default: "")
# WANDB_API_KEY W&B API key (default: "")
# WANDB_MODE W&B mode (default: online)
set -euo pipefail
CONFIG="${1:?Usage: $0 <config.yaml> <num_nodes> [extra flags...]}"
NUM_NODES="${2:?Usage: $0 <config.yaml> <num_nodes> [extra flags...]}"
shift 2
# ── Defaults ──────────────────────────────────────────────────────
PARTITION="${PARTITION:-main}"
NUM_GPUS="${NUM_GPUS:-8}"
CPUS_PER_TASK="${CPUS_PER_TASK:-128}"
MEM="${MEM:-1440G}"
MASTER_PORT="${MASTER_PORT:-29500}"
EXCLUDE="${EXCLUDE:-}"
WANDB_API_KEY="${WANDB_API_KEY:-}"
WANDB_MODE="${WANDB_MODE:-online}"
TOTAL_GPUS=$(( NUM_NODES * NUM_GPUS ))
CONFIG_NAME="$(basename "${CONFIG}" .yaml)"
JOB_NAME="${JOB_NAME:-${CONFIG_NAME}}"
OUTPUT_DIR="${OUTPUT_DIR:-logs/slurm}"
mkdir -p "${OUTPUT_DIR}"
# ── Build sbatch args ─────────────────────────────────────────────
SBATCH_ARGS=(
--job-name="${JOB_NAME}"
--partition="${PARTITION}"
--nodes="${NUM_NODES}"
--ntasks="${NUM_NODES}"
--ntasks-per-node=1
--gres="gpu:${NUM_GPUS}"
--cpus-per-task="${CPUS_PER_TASK}"
--mem="${MEM}"
--output="${OUTPUT_DIR}/${JOB_NAME}_%j.out"
--error="${OUTPUT_DIR}/${JOB_NAME}_%j.err"
--exclusive
)
if [[ -n "${EXCLUDE}" ]]; then
SBATCH_ARGS+=(--exclude="${EXCLUDE}")
fi
# ── Collect extra overrides for the training script ───────────────
EXTRA_ARGS=("$@")
echo "=== Slurm Training Submission ==="
echo "Config: ${CONFIG}"
echo "Nodes: ${NUM_NODES}"
echo "GPUs/node: ${NUM_GPUS}"
echo "Total GPUs: ${TOTAL_GPUS}"
echo "Partition: ${PARTITION}"
echo "Job name: ${JOB_NAME}"
echo "Extra args: ${EXTRA_ARGS[*]:-}"
echo "================================="
# ── Submit ────────────────────────────────────────────────────────
sbatch "${SBATCH_ARGS[@]}" <<EOF
#!/bin/bash
set -e -x
# ── Environment ───────────────────────────────────────────────────
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export TRITON_CACHE_DIR=/tmp/triton_cache_\${SLURM_PROCID}
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="${WANDB_API_KEY}"
export WANDB_MODE="${WANDB_MODE}"
# ── Rendezvous ────────────────────────────────────────────────────
export MASTER_PORT=${MASTER_PORT}
nodes=( \$(scontrol show hostnames \$SLURM_JOB_NODELIST) )
export MASTER_ADDR=\${nodes[0]}
export NODE_RANK=\$SLURM_PROCID
echo "MASTER_ADDR: \$MASTER_ADDR"
echo "NODE_RANK: \$NODE_RANK"
# ── Launch ────────────────────────────────────────────────────────
srun torchrun \\
--nnodes \$SLURM_JOB_NUM_NODES \\
--nproc_per_node ${NUM_GPUS} \\
--node_rank \$SLURM_PROCID \\
--rdzv_backend=c10d \\
--rdzv_endpoint="\$MASTER_ADDR:\$MASTER_PORT" \\
fastvideo/train/entrypoint/train.py \\
--config ${CONFIG} \\
--training.distributed.num_gpus ${TOTAL_GPUS} \\
${EXTRA_ARGS[*]:-}
EOF
@@ -2,35 +2,15 @@
"data": [
{
"caption": "In the video, a woman is elegantly showcasing her earrings, bringing attention to their intricate design with a gentle touch of her fingers. She is bathed in ambient purple and pink lighting, which casts a soft glow on her delicate features and enhances the vivid tones of her lipstick and eye makeup. Her hair is styled to frame her face smoothly, emphasizing the contours of her jawline and cheekbones. The background features a blurred neon light, adding an artistic and modern touch to the overall aesthetic.",
"video_path": "Fashion/mixkit-face-of-an-elegant-and-captivating-woman-41914_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
"video_path": "Man/mixkit-a-rancher-riding-a-horse-at-sunset-1143_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a dimly lit, eerie setting, a mysterious pink bottle labeled \"Authentic 100% organic POISON\" sits prominently in the foreground, casting a menacing aura. The bottle is accentuated by green fog, which swirls lightly around it, enhancing its sinister allure. Behind it, a shadowy golden bottle adorned with a spider emblem subtly emerges, adding an extra layer of mystery to the scene. Dim candles provide faint, flickering light, which complements the dark atmosphere, making the setting ideal for an illusion of hidden dangers.",
"video_path": "smoke/mixkit-poison-in-halloween-ritual-33879_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.",
"video_path": "forest/mixkit-view-of-a-river-between-two-old-trees-560_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
}
]
}
@@ -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
}
]
}
@@ -6,6 +6,7 @@ DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
--master_port=29513 \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
+9
View File
@@ -7,6 +7,13 @@ CUDA kernels for FastVideo video generation.
### Standard Installation (Local Development)
This will automatically detect your GPU architecture. If an NVIDIA Hopper (H100/sm_90a) GPU is detected, ThunderKittens kernels will be enabled. Otherwise, they will be skipped, and the package will use Triton fallbacks at runtime.
Before installation, set CUDA toolchain paths:
```bash
export CUDA_HOME=/usr/local/cuda
export CUDACXX=$CUDA_HOME/bin/nvcc
```
```bash
git submodule update --init --recursive
cd fastvideo-kernel
@@ -62,6 +69,8 @@ This package also includes kernels from [TurboDiffusion](https://github.com/thu-
- Any CUDA GPU for Triton-based fallbacks.
- **Build**:
- CUDA Toolkit 12.3+
- `CUDA_HOME` must be set (for example, `/usr/local/cuda`)
- `CUDACXX` must be set (for example, `$CUDA_HOME/bin/nvcc`)
- C++20 compatible compiler (GCC 10+, Clang 11+)
## Acknowledgement
+63 -8
View File
@@ -3,8 +3,10 @@ set -ex
# Simple build script wrapping uv/pip
# Usage:
# ./build.sh # local dev build (auto-detect / skip TK kernels when not available)
# ./build.sh --release # force-enable Hopper/TK kernels for release builds (no GPU required)
# ./build.sh # local build (torch-based arch detection, TK only on SM90)
# Environment overrides (if set, they win over auto-detection):
# TORCH_CUDA_ARCH_LIST
# CMAKE_ARGS (for FASTVIDEO_KERNEL_BUILD_TK / CMAKE_CUDA_ARCHITECTURES / GPU_BACKEND)
echo "Building fastvideo-kernel..."
@@ -12,7 +14,7 @@ echo "Building fastvideo-kernel..."
git submodule update --init --recursive
# Install build dependencies
pip install scikit-build-core cmake ninja
uv pip install scikit-build-core cmake ninja
RELEASE=0
GPU_BACKEND=CUDA
@@ -24,15 +26,68 @@ for arg in "$@"; do
esac
done
# Force-enable ThunderKittens kernels and compile for Hopper.
export TORCH_CUDA_ARCH_LIST="9.0a"
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
has_cmake_arg() {
local key="$1"
[[ "${CMAKE_ARGS:-}" =~ (^|[[:space:]])-D${key}(=|$) ]]
}
export CMAKE_ARGS="${CMAKE_ARGS:-} -DGPU_BACKEND=${GPU_BACKEND}"
detect_with_torch() {
uv run --active --no-project python -c "import torch
if not torch.cuda.is_available():
raise RuntimeError('torch.cuda.is_available() is false')
mj, mn = torch.cuda.get_device_capability(0)
print(f'{mj}.{mn}')"
}
if [ "${GPU_BACKEND}" = "CUDA" ]; then
detected_cc="$(detect_with_torch)" || {
echo "ERROR: torch-based CUDA arch detection failed in uv environment." >&2
echo " Ensure torch is installed and CUDA is available in the uv-selected Python." >&2
exit 1
}
cc_major="${detected_cc%%.*}"
cc_minor="${detected_cc##*.}"
cmake_arch="${cc_major}${cc_minor}"
echo "Detected compute capability via torch: ${detected_cc} (sm_${cmake_arch})"
# Respect explicit overrides.
if [ -z "${TORCH_CUDA_ARCH_LIST:-}" ]; then
if [ "${cc_major}" = "9" ] && [ "${cc_minor}" = "0" ]; then
export TORCH_CUDA_ARCH_LIST="9.0a"
else
export TORCH_CUDA_ARCH_LIST="${cc_major}.${cc_minor}"
fi
fi
# ThunderKittens build targeting:
# - SM90: compile Hopper/TK kernels with 90a.
# - Others (e.g., SM100): compile non-TK path with detected arch.
if ! has_cmake_arg "CMAKE_CUDA_ARCHITECTURES"; then
if [ "${cc_major}" = "9" ] && [ "${cc_minor}" = "0" ]; then
CMAKE_ARGS="${CMAKE_ARGS:-} -DCMAKE_CUDA_ARCHITECTURES=90a"
else
CMAKE_ARGS="${CMAKE_ARGS:-} -DCMAKE_CUDA_ARCHITECTURES=${cmake_arch}"
fi
fi
if ! has_cmake_arg "FASTVIDEO_KERNEL_BUILD_TK"; then
if [ "${cc_major}" = "9" ] && [ "${cc_minor}" = "0" ]; then
CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON"
else
CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=OFF"
fi
fi
fi
if ! has_cmake_arg "GPU_BACKEND"; then
CMAKE_ARGS="${CMAKE_ARGS:-} -DGPU_BACKEND=${GPU_BACKEND}"
fi
export CMAKE_ARGS
echo "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST:-<unset>}"
echo "CMAKE_ARGS: ${CMAKE_ARGS:-<unset>}"
echo "GPU_BACKEND: ${GPU_BACKEND:-<unset>}"
# Build and install
# Use -v for verbose output
pip install . -v --no-build-isolation
uv pip install . -v --no-build-isolation
@@ -1,20 +1,15 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
from .triton_kernels.index import map_to_index
# Try to load the C++ extension
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
sta_fwd = getattr(fastvideo_kernel_ops, "sta_fwd", None)
block_sparse_fwd = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
block_sparse_bwd = getattr(fastvideo_kernel_ops, "block_sparse_bwd", None)
except ImportError:
sta_fwd = None
block_sparse_fwd = None
block_sparse_bwd = None
def sliding_tile_attention(
q: torch.Tensor,
@@ -135,14 +130,8 @@ def video_sparse_attn(
mask = torch.zeros_like(scores,
dtype=torch.bool).scatter_(-1, topk_idx, True)
idx, num = map_to_index(mask)
if block_sparse_fwd is not None:
# Use autograd-enabled wrapper so backward works (and still uses SM90 kernel when available)
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
else:
# Triton-only forward (kept for environments without the wrapper deps)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1,405 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
from collections import defaultdict
from typing import Any
import numpy as np
from fastvideo.utils import dict_to_3d_list
def configure_sta(mode: str = 'STA_searching',
layer_num: int = 40,
time_step_num: int = 50,
head_num: int = 40,
**kwargs) -> list[list[list[Any]]]:
"""
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
Parameters:
----------
mode : str
The STA mode to use. Options are:
- 'STA_searching': Generate a set of mask candidates for initial search
- 'STA_tuning': Select best mask strategy based on previously saved results
- 'STA_inference': Load and use a previously tuned mask strategy
layer_num: int, number of layers
time_step_num: int, number of timesteps
head_num: int, number of heads
**kwargs : dict
Mode-specific parameters:
For 'STA_searching':
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
For 'STA_tuning':
- mask_search_files_path: str, required, path to mask search results
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
For 'STA_inference':
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
"""
valid_modes = [
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
]
if mode not in valid_modes:
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
if mode == 'STA_searching':
# Get parameters with defaults
mask_candidates: list[str] | None = kwargs.get('mask_candidates')
if mask_candidates is None:
raise ValueError(
"mask_candidates is required for STA_searching mode")
mask_selected: list[int] = kwargs.get('mask_selected',
list(range(len(mask_candidates))))
# Parse selected masks
selected_masks: list[list[int]] = []
for index in mask_selected:
mask = mask_candidates[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Create 3D mask structure with fixed dimensions (t=50, l=60)
masks_3d: list[list[list[list[int]]]] = []
for i in range(time_step_num): # Fixed t dimension = 50
row = []
for j in range(layer_num): # Fixed l dimension = 60
row.append(selected_masks) # Add all masks at each position
masks_3d.append(row)
return masks_3d
elif mode == 'STA_tuning':
# Get required parameters
mask_search_files_path: str | None = kwargs.get(
'mask_search_files_path')
if not mask_search_files_path:
raise ValueError(
"mask_search_files_path is required for STA_tuning mode")
# Get optional parameters with defaults
mask_candidates_tuning: list[str] | None = kwargs.get('mask_candidates')
if mask_candidates_tuning is None:
raise ValueError("mask_candidates is required for STA_tuning mode")
mask_selected_tuning: list[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_tuning))))
skip_time_steps_tuning: int | None = kwargs.get('skip_time_steps')
save_dir_tuning: str | None = kwargs.get('save_dir', "mask_candidates")
# Parse selected masks
selected_masks_tuning: list[list[int]] = []
for index in mask_selected_tuning:
mask = mask_candidates_tuning[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_tuning.append(masks_list)
# Read JSON results
results = read_specific_json_files(mask_search_files_path)
averaged_results = average_head_losses(results, selected_masks_tuning)
# Add full attention mask for specific cases
full_attention_mask_tuning: list[int] | None = kwargs.get(
'full_attention_mask')
if full_attention_mask_tuning is not None:
selected_masks_tuning.append(full_attention_mask_tuning)
# Select best mask strategy
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_tuning is None:
skip_time_steps_tuning = 12
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
timesteps_tuning, head_num)
# Save mask strategy
if save_dir_tuning is not None:
os.makedirs(save_dir_tuning, exist_ok=True)
file_path = os.path.join(
save_dir_tuning,
f'mask_strategy_s{skip_time_steps_tuning}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
elif mode == 'STA_tuning_cfg':
# Get required parameters for both positive and negative paths
mask_search_files_path_pos: str | None = kwargs.get(
'mask_search_files_path_pos')
mask_search_files_path_neg: str | None = kwargs.get(
'mask_search_files_path_neg')
save_dir_cfg: str | None = kwargs.get('save_dir')
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
raise ValueError(
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
)
# Get optional parameters with defaults
mask_candidates_cfg: list[str] | None = kwargs.get('mask_candidates')
if mask_candidates_cfg is None:
raise ValueError(
"mask_candidates is required for STA_tuning_cfg mode")
mask_selected_cfg: list[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_cfg))))
skip_time_steps_cfg: int | None = kwargs.get('skip_time_steps')
# Parse selected masks
selected_masks_cfg: list[list[int]] = []
for index in mask_selected_cfg:
mask = mask_candidates_cfg[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_cfg.append(masks_list)
# Read JSON results for both positive and negative paths
pos_results = read_specific_json_files(mask_search_files_path_pos)
neg_results = read_specific_json_files(mask_search_files_path_neg)
# Combine positive and negative results into one list
combined_results = pos_results + neg_results
# Average the combined results
averaged_results = average_head_losses(combined_results,
selected_masks_cfg)
# Add full attention mask for specific cases
full_attention_mask_cfg: list[int] | None = kwargs.get(
'full_attention_mask')
if full_attention_mask_cfg is not None:
selected_masks_cfg.append(full_attention_mask_cfg)
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_cfg is None:
skip_time_steps_cfg = 12
# Select best mask strategy using combined results
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
timesteps_cfg, head_num)
# Save mask strategy
os.makedirs(save_dir_cfg, exist_ok=True)
file_path = os.path.join(save_dir_cfg,
f'mask_strategy_s{skip_time_steps_cfg}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
else: # STA_inference
# Get parameters with defaults
load_path: str | None = kwargs.get(
'load_path', "mask_candidates/mask_strategy.json")
if load_path is None:
raise ValueError("load_path is required for STA_inference mode")
# Load previously saved mask strategy
with open(load_path) as f:
mask_strategy = json.load(f)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
# Helper functions
def read_specific_json_files(folder_path: str) -> list[dict[str, Any]]:
"""Read and parse JSON files containing mask search results."""
json_contents: list[dict[str, Any]] = []
# List files only in the current directory (no walk)
files = os.listdir(folder_path)
# Filter files
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
print(f"Found {len(matching_files)} matching files: {matching_files}")
for file_name in matching_files:
file_path = os.path.join(folder_path, file_name)
with open(file_path) as file:
data = json.load(file)
json_contents.append(data)
return json_contents
def average_head_losses(
results: list[dict[str, Any]],
selected_masks: list[list[int]]) -> dict[str, dict[str, np.ndarray]]:
"""Average losses across all prompts for each mask strategy."""
# Initialize a dictionary to store the averaged results
averaged_losses: dict[str, dict[str, np.ndarray]] = {}
loss_type = 'L2_loss'
# Get all loss types (e.g., 'L2_loss')
averaged_losses[loss_type] = {}
for mask in selected_masks:
mask_str = str(mask)
data_shape = np.array(results[0][loss_type][mask_str]).shape
accumulated_data = np.zeros(data_shape)
# Sum across all prompts
for prompt_result in results:
accumulated_data += np.array(prompt_result[loss_type][mask_str])
# Average by dividing by number of prompts
averaged_data = accumulated_data / len(results)
averaged_losses[loss_type][mask_str] = averaged_data
return averaged_losses
def select_best_mask_strategy(
averaged_results: dict[str, dict[str, np.ndarray]],
selected_masks: list[list[int]],
skip_time_steps: int = 12,
timesteps: int = 50,
head_num: int = 40
) -> tuple[dict[str, list[int]], float, dict[str, int]]:
"""Select the best mask strategy for each head based on loss minimization."""
best_mask_strategy: dict[str, list[int]] = {}
loss_type = 'L2_loss'
# Get the shape of time steps and layers
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
# Counter for sparsity calculation
total_tokens = 0 # total number of masked tokens
total_length = 0 # total sequence length
strategy_counts: dict[str, int] = {
str(strategy): 0
for strategy in selected_masks
}
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
for t in range(timesteps):
for layer_idx in range(layers):
for h in range(head_num):
if t < skip_time_steps: # First steps use full attention
strategy = full_attn_strategy
else:
# Get losses for this head across all strategies
head_losses = []
for strategy in selected_masks[:
-1]: # Exclude full attention
head_losses.append(averaged_results[loss_type][str(
strategy)][t][layer_idx][h])
# Find which strategy gives minimum loss
best_strategy_idx = np.argmin(head_losses)
strategy = selected_masks[best_strategy_idx]
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
# Calculate sparsity
nums = strategy # strategy is already a list of numbers
total_tokens += nums[0] * nums[1] * nums[
2] # masked tokens for chosen strategy
total_length += full_attn_strategy[0] * full_attn_strategy[
1] * full_attn_strategy[2]
# Count strategy usage
strategy_counts[str(strategy)] += 1
overall_sparsity = 1 - total_tokens / total_length
return best_mask_strategy, overall_sparsity, strategy_counts
def save_mask_search_results(
mask_search_final_result: list[dict[str, list[float]]],
prompt: str,
mask_strategies: list[str],
output_dir: str = 'output/mask_search_result/') -> str | None:
if not mask_search_final_result:
print("No mask search results to save")
return None
# Create result dictionary with defaultdict for nested lists
mask_search_dict: dict[str, dict[str, list[list[float]]]] = {
"L2_loss": defaultdict(list),
"L1_loss": defaultdict(list)
}
mask_selected = list(range(len(mask_strategies)))
selected_masks: list[list[int]] = []
for index in mask_selected:
mask = mask_strategies[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Process each mask strategy
for i, mask_strategy in enumerate(selected_masks):
mask_strategy_str = str(mask_strategy)
# Process L2 loss
step_results: list[list[float]] = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L2_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
step_results.append(layer_losses)
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
step_results = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L1_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
step_results.append(layer_losses)
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
# Create a filename based on the first 20 characters of the prompt
filename = prompt[:50].replace(" ", "_")
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
# Save the results to a JSON file
with open(filepath, 'w') as f:
json.dump(mask_search_dict, f, indent=4)
print(f"Successfully saved mask research results to {filepath}")
return filepath
+28 -15
View File
@@ -4,26 +4,36 @@ import torch
import torch.nn.functional as F
from flash_attn import flash_attn_func as flash_attn_2_func
from dataclasses import dataclass
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
# flash_attn 3 no longer have a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "4"
except ImportError:
flash_attn_func = flash_attn_2_func
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
# flash_attn 3 no longer have a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
flash_attn_func = flash_attn_2_func
fa_version = "2"
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
logger.info("Using FlashAttention-%s backend", fa_version)
class FlashAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
@@ -117,11 +127,13 @@ class FlashAttentionImpl(AttentionImpl):
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)
flash_attn_no_pad,
flash_attn_varlen_qk_no_pad,
)
attn_mask = attn_metadata.attn_mask
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
@@ -160,5 +172,6 @@ class FlashAttentionImpl(AttentionImpl):
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
causal=self.causal,
)
return output
@@ -1,263 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import dataclass
from typing import Any
import torch
from einops import rearrange
from fastvideo_kernel import sliding_tile_attention
import fastvideo.envs as envs
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.distributed import get_sp_group
from fastvideo.forward_context import ForwardContext, get_forward_context
from fastvideo.logger import init_logger
from fastvideo.utils import dict_to_3d_list
logger = init_logger(__name__)
class RangeDict(dict):
def __getitem__(self, item: int) -> str:
for key in self.keys():
if isinstance(key, tuple):
low, high = key
if low <= item <= high:
return str(super().__getitem__(key))
elif key == item:
return str(super().__getitem__(key))
raise KeyError(f"seq_len {item} not supported for STA")
class SlidingTileAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
# TODO(will-refactor): check this
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
return "SLIDING_TILE_ATTN"
@staticmethod
def get_impl_cls() -> type["SlidingTileAttentionImpl"]:
return SlidingTileAttentionImpl
@staticmethod
def get_metadata_cls() -> type["SlidingTileAttentionMetadata"]:
return SlidingTileAttentionMetadata
@staticmethod
def get_builder_cls() -> type["SlidingTileAttentionMetadataBuilder"]:
return SlidingTileAttentionMetadataBuilder
@dataclass
class SlidingTileAttentionMetadata(AttentionMetadata):
current_timestep: int
STA_param: list[list[
Any]] # each timestep with one metadata, shape [num_layers, num_heads]
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
STA_param: list[list[Any]],
current_timestep: int,
**kwargs: dict[str, Any],
) -> SlidingTileAttentionMetadata:
param = STA_param
if param is None:
return SlidingTileAttentionMetadata(
current_timestep=current_timestep, STA_param=[])
return SlidingTileAttentionMetadata(current_timestep=current_timestep,
STA_param=param[current_timestep])
class SlidingTileAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
# TODO(will-refactor): for now this is the mask strategy, but maybe we should
# have a more general config for STA?
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
# TODO(kevin): get mask strategy for different STA modes
with open(config_file) as f:
mask_strategy = json.load(f)
self.mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
# STA config
self.STA_base_tile_size = [6, 8, 8]
self.dit_seq_shape_mapping = RangeDict({
(115200, 115456): '30x48x80',
82944: '36x48x48',
69120: '18x48x80',
})
self.full_window_mapping = {
'30x48x80': [5, 6, 10],
'36x48x48': [6, 6, 6],
'18x48x80': [3, 6, 10]
}
def tile(self, x: torch.Tensor) -> torch.Tensor:
return rearrange(
x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
def untile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(
x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
return x
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
img_sequence_length = qkv.shape[1]
self.dit_seq_shape_str = self.dit_seq_shape_mapping[img_sequence_length]
self.full_window_size = self.full_window_mapping[self.dit_seq_shape_str]
self.dit_seq_shape_int = list(
map(int, self.dit_seq_shape_str.split('x')))
self.img_seq_length = self.dit_seq_shape_int[
0] * self.dit_seq_shape_int[1] * self.dit_seq_shape_int[2]
return self.tile(qkv)
def postprocess_output(
self,
output: torch.Tensor,
attn_metadata: SlidingTileAttentionMetadata,
) -> torch.Tensor:
return self.untile(output)
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
attn_metadata: SlidingTileAttentionMetadata,
) -> torch.Tensor:
if self.mask_strategy is None:
raise ValueError(
"mask_strategy cannot be None for SlidingTileAttention")
if self.mask_strategy[0] is None:
raise ValueError(
"mask_strategy[0] cannot be None for SlidingTileAttention")
timestep = attn_metadata.current_timestep
forward_context: ForwardContext = get_forward_context()
forward_batch = forward_context.forward_batch
if forward_batch is None:
raise ValueError("forward_batch cannot be None")
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
if attn_metadata.STA_param is None or len(
attn_metadata.STA_param) <= layer_idx:
raise ValueError("Invalid STA_param")
STA_param = attn_metadata.STA_param[layer_idx]
text_length = q.shape[1] - self.img_seq_length
has_text = text_length > 0
query = q.transpose(1, 2).contiguous()
key = k.transpose(1, 2).contiguous()
value = v.transpose(1, 2).contiguous()
head_num = query.size(1)
sp_group = get_sp_group()
current_rank = sp_group.rank_in_group
start_head = current_rank * head_num
# searching or tuning mode
if len(STA_param) < head_num * sp_group.world_size:
sparse_attn_hidden_states_all = []
full_mask_window = STA_param[-1]
for window_size in STA_param[:-1]:
sparse_hidden_states = sliding_tile_attention(
query, key, value, [window_size] * head_num, text_length,
has_text, self.dit_seq_shape_str).transpose(1, 2)
sparse_attn_hidden_states_all.append(sparse_hidden_states)
hidden_states = sliding_tile_attention(
query, key, value, [full_mask_window] * head_num, text_length,
has_text, self.dit_seq_shape_str).transpose(1, 2)
attn_L2_loss = []
attn_L1_loss = []
# average loss across all heads
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
# L2 loss
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() -
hidden_states.float())**2,
dim=[0, 1, 3]).cpu().numpy()
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
attn_L2_loss.append(attn_L2_loss_)
# L1 loss
attn_L1_loss_ = torch.mean(
torch.abs(sparse_attn_hidden_states.float() -
hidden_states.float()),
dim=[0, 1, 3]).cpu().numpy()
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
attn_L1_loss.append(attn_L1_loss_)
layer_loss_save = {"L2_loss": attn_L2_loss, "L1_loss": attn_L1_loss}
if forward_batch.is_cfg_negative:
if forward_batch.mask_search_final_result_neg is not None:
forward_batch.mask_search_final_result_neg[timestep].append(
layer_loss_save)
else:
if forward_batch.mask_search_final_result_pos is not None:
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.dit_seq_shape_str).transpose(1, 2)
return hidden_states
+15 -26
View File
@@ -66,11 +66,11 @@ class DistributedAttention(nn.Module):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
original_seq_len: int | None = None,
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
@@ -78,10 +78,10 @@ class DistributedAttention(nn.Module):
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
original_seq_len (int): Original (unpadded) full sequence length
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
@@ -91,7 +91,7 @@ class DistributedAttention(nn.Module):
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim(
) == 4, "Expected 4D tensors"
batch_size, seq_len, num_heads, head_dim = q.shape
batch_size, _, num_heads, _ = q.shape
local_rank = get_sp_parallel_rank()
world_size = get_sp_world_size()
@@ -107,15 +107,11 @@ class DistributedAttention(nn.Module):
scatter_dim=2,
gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads
# The attention mask should now apply to the full sequence length
# Since mask is [batch, full_seq_len], it's already in the correct format
# LOAY TODO, instead of slicing repeatedly maintain an original qkv and rewrite into that
valid_seq_len = None
if attention_mask is not None:
valid_seq_len = (attention_mask[0] == 1).sum().item()
qkv = qkv[:, :valid_seq_len, :, :]
# 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.
original_seq_len = original_seq_len or qkv.shape[1]
pad_seq_len = qkv.shape[1] - original_seq_len
qkv = qkv[:, :original_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
@@ -145,7 +141,7 @@ class DistributedAttention(nn.Module):
# Redistribute back if using sequence parallelism
replicated_output = None
if replicated_q is not None:
split_idx = seq_len * world_size if valid_seq_len is None else valid_seq_len
split_idx = original_seq_len
replicated_output = output[:, split_idx:]
output = output[:, :split_idx]
# TODO: make this asynchronous
@@ -154,9 +150,7 @@ class DistributedAttention(nn.Module):
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
if attention_mask is not None:
pad_len = (attention_mask[0] == 0).sum().item()
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
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,
@@ -175,12 +169,12 @@ class DistributedAttention_VSA(DistributedAttention):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
original_seq_len: int,
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
gate_compress: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
@@ -188,11 +182,11 @@ class DistributedAttention_VSA(DistributedAttention):
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
original_seq_len (int): Original (unpadded) full sequence length
gate_compress (torch.Tensor): Gate compress tensor [batch_size, seq_len, num_heads, head_dim]
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
@@ -221,11 +215,8 @@ class DistributedAttention_VSA(DistributedAttention):
gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads
# The attention mask should now apply to the full sequence length
if attention_mask is not None:
valid_seq_len = (attention_mask[0] == 1).sum().item()
qkvg = qkvg[:, :valid_seq_len, :, :]
pad_seq_len = qkvg.shape[1] - original_seq_len
qkvg = qkvg[:, :original_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
@@ -246,9 +237,7 @@ class DistributedAttention_VSA(DistributedAttention):
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
if attention_mask is not None:
pad_len = (attention_mask[0] == 0).sum().item()
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
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,
@@ -0,0 +1,264 @@
from __future__ import annotations
import torch
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"
)
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})"
)
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
out, lse = _flash_attn_fwd(
q,
k,
v,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=1,
pack_gqa=None,
)
return out, lse
@torch.library.register_fake("fastvideo::_flash_attn_cute_forward")
def _flash_attn_cute_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del k, softmax_scale, causal, deterministic
batch, seqlen_q, nheads = q.shape[:3]
out = q.new_empty(batch, seqlen_q, nheads, v.shape[-1])
lse = q.new_empty(batch, nheads, seqlen_q, dtype=torch.float32)
return out, lse
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)
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.deterministic = deterministic
def _flash_attn_cute_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
deterministic=ctx.deterministic,
)
return dq, dk, dv, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_forward",
_flash_attn_cute_backward,
setup_context=_flash_attn_cute_setup_context,
)
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_varlen_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_varlen_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
out, lse = _flash_attn_fwd(
q,
k,
v,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=1,
pack_gqa=None,
)
return out, lse
@torch.library.register_fake("fastvideo::_flash_attn_cute_varlen_forward")
def _flash_attn_cute_varlen_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del k, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, softmax_scale
del causal
del deterministic
total_q, nheads = q.shape[:2]
out = q.new_empty(total_q, nheads, v.shape[-1])
lse = q.new_empty(nheads, total_q, dtype=torch.float32)
return out, lse
def _flash_attn_cute_varlen_setup_context(
ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
softmax_scale,
causal,
deterministic,
) = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k)
ctx.max_seqlen_q = max_seqlen_q
ctx.max_seqlen_k = max_seqlen_k
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.deterministic = deterministic
def _flash_attn_cute_varlen_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=ctx.max_seqlen_q,
max_seqlen_k=ctx.max_seqlen_k,
deterministic=ctx.deterministic,
)
return dq, dk, dv, None, None, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_varlen_forward",
_flash_attn_cute_varlen_backward,
setup_context=_flash_attn_cute_varlen_setup_context,
)
def flash_attn_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
causal: bool = False,
deterministic: bool = False,
) -> 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)
return out
def flash_attn_varlen_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
causal: bool = False,
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
_check_dropout(dropout_p)
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
softmax_scale,
causal,
deterministic,
)
return out
+65 -45
View File
@@ -15,16 +15,29 @@
# See the License for the specific language governing permissions and limitations under the License.
from einops import rearrange
from flash_attn.bert_padding import pad_input, unpad_input
from flash_attn import flash_attn_varlen_qkvpacked_func
try:
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_impl, )
except ImportError:
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_impl, )
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_impl, )
def flash_attn_no_pad(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
def flash_attn_no_pad(
qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False,
):
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
@@ -46,22 +59,28 @@ def flash_attn_no_pad(qkv,
deterministic=deterministic,
)
output = rearrange(
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices,
batch_size, seqlen),
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
indices,
batch_size,
seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
def flash_attn_no_pad_v3(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
from flash_attn.bert_padding import pad_input, unpad_input
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_v3
def flash_attn_no_pad_v3(
qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False,
):
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_v3, )
if flash_attn_varlen_func_v3 is None:
raise ImportError("FlashAttention V3 backend not available")
@@ -80,22 +99,29 @@ def flash_attn_no_pad_v3(qkv,
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
output_unpad = flash_attn_varlen_func_v3(query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_q,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic)
output_unpad = flash_attn_varlen_func_v3(
query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_q,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size,
seqlen),
"b s (h d) -> b s h d",
h=nheads)
output = rearrange(
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
indices,
batch_size,
seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
@@ -110,16 +136,6 @@ def flash_attn_varlen_qk_no_pad(
softmax_scale=None,
deterministic=False,
):
from flash_attn.bert_padding import pad_input, unpad_input
try:
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_impl
except ImportError:
from flash_attn import flash_attn_varlen_func as flash_attn_varlen_func_impl
if flash_attn_varlen_func_impl is None:
raise ImportError("FlashAttention varlen backend not available")
batch_size, q_seqlen, nheads, _ = query.shape
query_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
@@ -148,8 +164,12 @@ def flash_attn_varlen_qk_no_pad(
)
output = rearrange(
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), q_indices,
batch_size, q_seqlen),
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
q_indices,
batch_size,
q_seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
+1 -2
View File
@@ -43,6 +43,5 @@
"require_post_norm": null
}
],
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
}
-1
View File
@@ -24,7 +24,6 @@ class ModelConfig:
arch_config: ArchConfig = field(default_factory=ArchConfig)
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
def __getattr__(self, name):
# Only called if 'name' is not found in ModelConfig directly
+3 -1
View File
@@ -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.wangamevideo import WanGameVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig",
"WanVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig",
"WanGameVideoConfig"
]
+4 -5
View File
@@ -15,11 +15,10 @@ class DiTArchConfig(ArchConfig):
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, 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)
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
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.wanvideo import (
WanVideoArchConfig,
WanVideoConfig,
)
@dataclass
class WanGameVideoArchConfig(WanVideoArchConfig):
"""Wangame keeps WanVideo architecture defaults and checkpoint mappings."""
@dataclass
class WanGameVideoConfig(WanVideoConfig):
arch_config: WanGameVideoArchConfig = field(
default_factory=WanGameVideoArchConfig
)
prefix: str = "WanGame"
@@ -78,7 +78,6 @@ class Qwen2_5_VLArchConfig(TextEncoderArchConfig):
"add_generation_prompt": True,
"tokenize": True,
"return_dict": True,
"padding": "max_length",
"max_length": 1000 + 108,
"truncation": True,
"return_tensors": "pt",
-1
View File
@@ -63,7 +63,6 @@ class T5ArchConfig(TextEncoderArchConfig):
self.dense_act_fn = "gelu_new"
self.tokenizer_kwargs = {
"padding": "max_length",
"truncation": True,
"max_length": self.text_len,
"add_special_tokens": True,

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