Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9d95b1c10 | ||
|
|
8b55e9706c |
@@ -1,94 +0,0 @@
|
||||
# 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.
|
||||
@@ -1,46 +0,0 @@
|
||||
# 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/`.
|
||||
@@ -1,48 +0,0 @@
|
||||
# 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.
|
||||
@@ -1,129 +0,0 @@
|
||||
# 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
|
||||
```
|
||||
@@ -1,327 +0,0 @@
|
||||
# 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
|
||||
@@ -1,21 +0,0 @@
|
||||
# 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`
|
||||
|
||||
-->
|
||||
@@ -1,4 +0,0 @@
|
||||
{"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"}
|
||||
@@ -1,34 +0,0 @@
|
||||
# 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._
|
||||
@@ -1,76 +0,0 @@
|
||||
# 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
|
||||
```
|
||||
@@ -1,302 +0,0 @@
|
||||
# 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.
|
||||
@@ -1,57 +0,0 @@
|
||||
---
|
||||
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"}
|
||||
```
|
||||
@@ -1,128 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,7 +0,0 @@
|
||||
{"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"}
|
||||
@@ -1,127 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,87 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,134 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,82 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,137 +0,0 @@
|
||||
---
|
||||
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 |
|
||||
@@ -1,54 +0,0 @@
|
||||
---
|
||||
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.
|
||||
@@ -1,47 +0,0 @@
|
||||
---
|
||||
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.
|
||||
@@ -1,87 +0,0 @@
|
||||
---
|
||||
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.
|
||||
@@ -1,71 +0,0 @@
|
||||
---
|
||||
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.
|
||||
@@ -1,67 +0,0 @@
|
||||
---
|
||||
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.
|
||||
@@ -1,46 +0,0 @@
|
||||
{
|
||||
"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
|
||||
}
|
||||
}
|
||||
}
|
||||
+60
-51
@@ -22,7 +22,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- TEST_TYPE=encoder
|
||||
@@ -35,7 +35,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- TEST_TYPE=vae
|
||||
@@ -61,7 +61,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 45m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
@@ -76,7 +76,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 20m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Inference Tests"
|
||||
env:
|
||||
- TEST_TYPE=inference_lora
|
||||
@@ -129,7 +129,11 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "fastvideo-kernel/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -140,16 +144,63 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo-kernel/**"
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Kernel Tests"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=kernel_tests
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo-kernel/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/tests/test_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/vmoba/**"
|
||||
- "fastvideo/attention/backends/vmoba.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
@@ -171,45 +222,3 @@ 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"
|
||||
# - "docker/Dockerfile.python3.12"
|
||||
# config:
|
||||
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
# label: "LoRA Extraction Tests"
|
||||
# env:
|
||||
# - TEST_TYPE=lora_extraction
|
||||
# agents:
|
||||
# queue: "default"
|
||||
|
||||
@@ -51,7 +51,6 @@ 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"
|
||||
@@ -76,7 +75,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_SSIM_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
@@ -90,9 +89,17 @@ 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"
|
||||
;;
|
||||
"kernel_tests")
|
||||
log "Running kernel tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
|
||||
"inference_sta")
|
||||
log "Running inference STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
|
||||
;;
|
||||
"precision_sta")
|
||||
log "Running precision STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
|
||||
;;
|
||||
"precision_vsa")
|
||||
log "Running precision VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
|
||||
;;
|
||||
"inference_lora")
|
||||
log "Running LoRA tests..."
|
||||
@@ -111,22 +118,14 @@ case "$TEST_TYPE" in
|
||||
log "Running V-MoBA inference tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
|
||||
;;
|
||||
"precision_vmoba")
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
"unit_test")
|
||||
log "Running unit tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
|
||||
;;
|
||||
"lora_extraction")
|
||||
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
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
## 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
|
||||
@@ -45,12 +45,6 @@ jobs:
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
- name: Generate docs examples
|
||||
run: python docs/generate_examples.py
|
||||
|
||||
- name: Check docs links
|
||||
run: python scripts/check_docs_links.py
|
||||
|
||||
- name: Build documentation
|
||||
run: mkdocs build
|
||||
|
||||
@@ -69,4 +63,4 @@ jobs:
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
uses: actions/deploy-pages@v4
|
||||
uses: actions/deploy-pages@v4
|
||||
@@ -1,225 +0,0 @@
|
||||
name: Publish FastVideo Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "fastvideo-kernel/pyproject.toml"
|
||||
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 fastvideo-kernel
|
||||
# Get current commit's version from pyproject.toml
|
||||
# Use ^ to match start of line to avoid matching minimum-version
|
||||
NEW_VERSION=$(grep -oP '^version\s*=\s*"\K[^"]+' pyproject.toml)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
# Note: git show expects path relative to repo root
|
||||
OLD_VERSION=$(git show HEAD~1:fastvideo-kernel/pyproject.toml | 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:
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12']
|
||||
torch-cuda:
|
||||
# - torch-version: '2.5.1'
|
||||
# cuda-version: '12.4.1'
|
||||
# torch-cuda-short: 'cu124'
|
||||
# - torch-version: '2.6.0'
|
||||
# cuda-version: '12.6.3'
|
||||
# torch-cuda-short: 'cu126'
|
||||
# - torch-version: '2.7.1'
|
||||
# cuda-version: '12.8.0'
|
||||
# torch-cuda-short: 'cu128'
|
||||
# - torch-version: '2.9.1'
|
||||
# cuda-version: '12.8.0'
|
||||
# torch-cuda-short: 'cu128'
|
||||
- torch-version: '2.10.0'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
cd fastvideo-kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
|
||||
|
||||
# Build standard wheel (no local version suffix) for PyPI
|
||||
python -m build --wheel --outdir dist
|
||||
|
||||
# Fix the wheel to be manylinux compliant
|
||||
pip install auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
import torch
|
||||
|
||||
print(os.path.join(os.path.dirname(torch.__file__), "lib"))
|
||||
PY
|
||||
)
|
||||
export LD_LIBRARY_PATH="${TORCH_LIB_DIR}:${LD_LIBRARY_PATH}"
|
||||
# Target manylinux_2_35 (Ubuntu 22.04 native)
|
||||
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist \
|
||||
--exclude libtorch_cuda.so \
|
||||
--exclude libtorch_cpu.so \
|
||||
--exclude libtorch.so \
|
||||
--exclude libc10.so \
|
||||
--exclude libc10_cuda.so \
|
||||
--exclude libtorch_python.so
|
||||
# Move fixed wheels back to dist for upload consistency
|
||||
rm dist/*.whl
|
||||
mv fixed_dist/*.whl dist/
|
||||
|
||||
- name: Upload wheel artifact
|
||||
# Only upload if it's the "main" CUDA version we want on PyPI
|
||||
# We upload all to artifacts for inspection/GH releases, but give them distinct artifact names
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: fastvideo_kernel-py${{ matrix.python-version }}-${{ matrix.torch-cuda.torch-cuda-short }}-torch${{ matrix.torch-cuda.torch-version }}
|
||||
path: fastvideo-kernel/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: Download PyPI wheels
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
path: fastvideo-kernel/dist/
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
pip install build scikit-build-core cmake ninja
|
||||
|
||||
cd fastvideo-kernel
|
||||
# We don't need full CUDA/Torch to just package the source (sdist)
|
||||
python -m build --sdist --outdir dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: fastvideo-kernel/dist/
|
||||
@@ -47,6 +47,16 @@ 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
|
||||
@@ -80,6 +90,8 @@ 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:
|
||||
@@ -94,6 +106,12 @@ 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/**'
|
||||
@@ -130,6 +148,13 @@ 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
|
||||
@@ -257,6 +282,44 @@ 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: >-
|
||||
@@ -315,7 +378,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, precision-test-VSA]
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, 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:
|
||||
@@ -332,7 +395,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", "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", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
@@ -0,0 +1,249 @@
|
||||
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/
|
||||
-17
@@ -18,7 +18,6 @@ venv/
|
||||
.venv/
|
||||
runs/
|
||||
samples/
|
||||
Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
@@ -31,13 +30,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
@@ -75,15 +67,6 @@ docs/distillation/examples/
|
||||
!docs/assets/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
!assets/images/**/*.png
|
||||
!assets/images/**/*.jpg
|
||||
!assets/images/**/*.jpeg
|
||||
!assets/images/**/*.gif
|
||||
!assets/videos/**/*.mp4
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
.claude/
|
||||
.codex/
|
||||
openspec/
|
||||
|
||||
+6
-5
@@ -1,6 +1,7 @@
|
||||
[submodule "fastvideo-kernel/include/tk"]
|
||||
path = fastvideo-kernel/include/tk
|
||||
[submodule "csrc/attn/video_sparse_attn/tk"]
|
||||
path = csrc/attn/video_sparse_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
[submodule "csrc/attn/sliding_tile_attn/tk"]
|
||||
path = csrc/attn/sliding_tile_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
[submodule "fastvideo-kernel/include/cutlass"]
|
||||
path = fastvideo-kernel/include/cutlass
|
||||
url = https://github.com/NVIDIA/cutlass.git
|
||||
|
||||
@@ -4,13 +4,13 @@ default_stages:
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/third_party/.*|
|
||||
fastvideo-kernel/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
assets/prompts/.*|
|
||||
prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
@@ -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-aiofiles, types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: check-filenames
|
||||
@@ -68,7 +68,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^\"*fastvideo/tests/ssim/" | grep -v "^\"*fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^fastvideo/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
|
||||
@@ -1,56 +0,0 @@
|
||||
# Repository Guidelines
|
||||
|
||||
## Project Structure & Module Organization
|
||||
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
|
||||
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
|
||||
- Tests:
|
||||
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
|
||||
- `tests/local_tests/` for additional local/component checks.
|
||||
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
|
||||
- Runnable examples and scripts: `examples/` and `scripts/`.
|
||||
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
- `pytest fastvideo/tests/ -v`: run package tests.
|
||||
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
|
||||
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
|
||||
|
||||
## Coding Style & Naming Conventions
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Target line length is 80.
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
|
||||
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
|
||||
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
|
||||
- Document GPU assumptions in tests that require specific hardware.
|
||||
|
||||
## Commit & Pull Request Guidelines
|
||||
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
|
||||
- Keep commits focused by concern (feature, refactor, fix).
|
||||
- PRs should include:
|
||||
- clear problem/solution summary,
|
||||
- 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.
|
||||
@@ -1,78 +1,70 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/fastwan.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- `2025/08/04`: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
### More News
|
||||
|
||||
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- End-to-end post-training support for bidirectional and autoregressive models:
|
||||
- End-to-end post-training support:
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
|
||||
- Data preprocessing pipeline for video data
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Data preprocessing pipeline for video, image, and text data
|
||||
- Distribution Matching Distillation (DMD2) stepwise distillation.
|
||||
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
|
||||
- Causal distillation through Self-Forcing
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- Sequence Parallelism for distributed inference
|
||||
- Multiple state-of-the-art attention backends
|
||||
- User-friendly CLI and Python API
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/optimizations/) for full list of supported optimizations.
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Diverse hardware and OS support
|
||||
- Support H100, A100, 4090
|
||||
- Support Linux, Windows, MacOS
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/support_matrix/) for full list of supported models, hardware assumptions, and optimization compatibility.
|
||||
|
||||
## Getting Started
|
||||
|
||||
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.
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
|
||||
```bash
|
||||
# Create and activate a new uv environment
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
|
||||
## Sparse Distillation
|
||||
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
| ------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
|:-------------------------------------------------------------------------------------------: |:---------------------------------------------------------------------------------------------------------------: |:--------------------------------------------------------------------------------------------------------: |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
|
||||
### Generating Your First Video
|
||||
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/attention/vsa/#installation). Create a file called `example.py` with the following code:
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
import os
|
||||
@@ -93,6 +85,7 @@ 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
|
||||
)
|
||||
@@ -109,37 +102,55 @@ python example.py
|
||||
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
|
||||
|
||||
## More Guides
|
||||
### Other docs:
|
||||
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
|
||||
## Awesome work using FastVideo or our research projects
|
||||
|
||||
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025.
|
||||
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo.
|
||||
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo.
|
||||
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo.
|
||||
- [HY-WorldPlay](https://github.com/Tencent-Hunyuan/HY-WorldPlay): An action-conditioned world model model trained using FastVideo framework.
|
||||
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention.
|
||||
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch.
|
||||
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention.
|
||||
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025. [](https://github.com/sgl-project/sglang)
|
||||
|
||||
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [](https://github.com/XueZeyue/DanceGRPO)
|
||||
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [](https://github.com/Tencent-Hunyuan/SRPO)
|
||||
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [](https://github.com/Vchitect/DCM)
|
||||
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
|
||||
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [](https://github.com/kandinskylab/kandinsky-5)
|
||||
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [](https://github.com/meituan-longcat/LongCat-Video)
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
|
||||
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
|
||||
- [Triton](https://github.com/triton-lang/triton)
|
||||
- [DMD2](https://github.com/tianweiy/DMD2)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
|
||||
If you find FastVideo useful, please consider citing our research work:
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
|
||||
```bibtex
|
||||
@software{fastvideo2024,
|
||||
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
|
||||
author = {The FastVideo Team},
|
||||
url = {https://github.com/hao-ai-lab/FastVideo},
|
||||
month = apr,
|
||||
year = {2024},
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={Vsa: Faster video diffusion with trainable sparse attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 490 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 1.2 MiB |
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@@ -1,7 +0,0 @@
|
||||
# FastVideo/assets/videos
|
||||
|
||||
This folder is used to store **video assets for examples**, primarily **input videos** consumed by scripts under `FastVideo/examples/`.
|
||||
|
||||
- **Typical contents**: short input clips for demos (e.g., video2world / image2video examples).
|
||||
- **Non-critical**: these assets are for convenience and are not required to use the FastVideo library.
|
||||
- **Large files**: avoid committing large videos to git; prefer shared storage or download-on-demand.
|
||||
Binary file not shown.
@@ -1,106 +0,0 @@
|
||||
# FVD (Fréchet Video Distance) Benchmark
|
||||
|
||||
Evaluate generated video quality using FVD with the I3D feature extractor.
|
||||
|
||||
## Quick Start
|
||||
|
||||
**Run the benchmark:**
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/run.sh
|
||||
```
|
||||
|
||||
That's it! The script auto-installs dependencies and runs the benchmark.
|
||||
|
||||
**To customize:** Edit `benchmarks/fvd/run_fvd.py` to change:
|
||||
- Video paths (`real_dir`, `gen_dir`)
|
||||
- Number of videos, frames, sampling strategy
|
||||
- Device, batch size, caching, etc.
|
||||
|
||||
## Advanced Usage (CLI)
|
||||
|
||||
For more control without editing Python files, use the CLI.
|
||||
|
||||
**First-time setup** (one-time per pod/environment):
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/setup_fvd.sh
|
||||
```
|
||||
|
||||
Then run any configuration you want:
|
||||
|
||||
```bash
|
||||
# Custom configuration
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--num-videos 1024 \
|
||||
--num-frames 32 \
|
||||
--clip-strategy random \
|
||||
--batch-size 32 \
|
||||
--seed 42 \
|
||||
--extractor clip
|
||||
```
|
||||
|
||||
**Standard protocols:**
|
||||
|
||||
```bash
|
||||
# Use predefined protocols
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
|
||||
```
|
||||
|
||||
This would use i3d model by default as the feature extractor
|
||||
|
||||
**Feature caching** (speed up repeated evaluations):
|
||||
|
||||
```bash
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f \
|
||||
--cache-real-features fvd-cache/extractor_name # Directory path (will save/load fvd-cache/extractor_name/extractor-name_real_features.pkl)
|
||||
```
|
||||
|
||||
Run `python -m benchmarks.fvd.cli --help` for all options.
|
||||
|
||||
## Available Protocols
|
||||
|
||||
- `fvd2048_16f` - Standard (2048 videos, 16 frames)
|
||||
- `fvd2048_128f` - Long videos (128 frames)
|
||||
- `fvd2048_128f_subsample8` - Subsampled long videos
|
||||
- `quick_test` - Fast testing (10 videos)
|
||||
|
||||
## Configuration Options
|
||||
|
||||
Key options in `FVDConfig`:
|
||||
|
||||
```python
|
||||
num_videos=2048, # Videos to evaluate
|
||||
num_frames_per_clip=16, # Frames per clip
|
||||
clip_strategy='beginning', # beginning|random|uniform|middle|sliding
|
||||
frame_stride=1, # Frame subsampling
|
||||
batch_size=32, # GPU batch size
|
||||
device='cuda', # cuda|cpu
|
||||
cache_real_features=None, # Cache path for speed
|
||||
seed=42, # Reproducibility
|
||||
extractor='i3d', # i3d|clip|videomae
|
||||
```
|
||||
|
||||
## Programmatic Usage
|
||||
|
||||
```python
|
||||
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
config = FVDConfig.fvd2048_16f() # or custom config
|
||||
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
print(f"FVD: {results['fvd']:.2f}")
|
||||
```
|
||||
|
||||
## Notes
|
||||
|
||||
- Requires minimum 10 frames per clip
|
||||
- Supports both video files (.mp4, .avi, etc.) and frame directories
|
||||
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
|
||||
@@ -1,38 +0,0 @@
|
||||
"""
|
||||
FastVideo Frechet Video Distance (FVD) Benchmark Module.
|
||||
>>> from fastvideo.benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
>>> config = FVDConfig.fvd2048_16f() # Standard protocol
|
||||
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
>>> print(f"FVD: {results['fvd']:.2f}")
|
||||
"""
|
||||
|
||||
from .fvd import (
|
||||
compute_fvd,
|
||||
compute_fvd_with_config,
|
||||
compute_frechet_distance,
|
||||
compute_statistics,
|
||||
FVDConfig,
|
||||
)
|
||||
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor,
|
||||
load_extractor)
|
||||
from .video_utils import (
|
||||
load_video_auto,
|
||||
sample_clips_from_video,
|
||||
load_video_clips_streaming,
|
||||
ClipSamplingStrategy,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'compute_fvd',
|
||||
'compute_fvd_with_config',
|
||||
'compute_frechet_distance',
|
||||
'compute_statistics',
|
||||
'FVDConfig',
|
||||
'BaseFeatureExtractor',
|
||||
'I3DFeatureExtractor',
|
||||
'load_extractor',
|
||||
'load_video_auto',
|
||||
'sample_clips_from_video',
|
||||
'load_video_clips_streaming',
|
||||
'ClipSamplingStrategy',
|
||||
]
|
||||
@@ -1,107 +0,0 @@
|
||||
import argparse
|
||||
import sys
|
||||
import traceback
|
||||
from .fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Compute Fréchet Video Distance (FVD)')
|
||||
|
||||
# Required arguments
|
||||
parser.add_argument('--real-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to real videos')
|
||||
parser.add_argument('--gen-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to generated videos')
|
||||
|
||||
# Extractor selection
|
||||
parser.add_argument('--extractor',
|
||||
type=str,
|
||||
default='i3d',
|
||||
choices=['i3d', 'clip', 'videomae'],
|
||||
help='Feature extractor model to use (default: i3d)')
|
||||
|
||||
# Standard args
|
||||
parser.add_argument('--seed',
|
||||
type=int,
|
||||
default=None,
|
||||
help='Random seed for reproducibility')
|
||||
parser.add_argument('--protocol',
|
||||
type=str,
|
||||
default=None,
|
||||
choices=['fvd2048_16f', 'fvd2048_128f', 'quick_test'],
|
||||
help='Use standard protocol (overrides other settings)')
|
||||
parser.add_argument('--num-videos',
|
||||
type=int,
|
||||
default=2048,
|
||||
help='Number of videos to use')
|
||||
parser.add_argument('--num-frames',
|
||||
type=int,
|
||||
default=16,
|
||||
help='Number of frames per clip')
|
||||
parser.add_argument('--clip-strategy',
|
||||
type=str,
|
||||
default='beginning',
|
||||
help='Clip sampling strategy')
|
||||
parser.add_argument('--batch-size',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Batch size for feature extraction')
|
||||
parser.add_argument('--device',
|
||||
type=str,
|
||||
default='cuda',
|
||||
help='Device to use (cuda or cpu)')
|
||||
parser.add_argument('--cache-real-features',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Path to cache real video features')
|
||||
parser.add_argument('--quiet',
|
||||
action='store_true',
|
||||
help='Suppress progress output')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create config
|
||||
if args.protocol:
|
||||
protocol_map = {
|
||||
'fvd2048_16f': FVDConfig.fvd2048_16f,
|
||||
'fvd2048_128f': FVDConfig.fvd2048_128f,
|
||||
'quick_test': FVDConfig.quick_test,
|
||||
}
|
||||
config = protocol_map[args.protocol]()
|
||||
# Apply overrides
|
||||
config.device = args.device
|
||||
config.cache_real_features = args.cache_real_features
|
||||
config.extractor_model = args.extractor # Apply extractor arg
|
||||
else:
|
||||
config = FVDConfig(
|
||||
num_videos=args.num_videos,
|
||||
num_frames_per_clip=args.num_frames,
|
||||
extractor_model=args.extractor, # Apply extractor arg
|
||||
clip_strategy=args.clip_strategy,
|
||||
batch_size=args.batch_size,
|
||||
device=args.device,
|
||||
cache_real_features=args.cache_real_features,
|
||||
seed=args.seed)
|
||||
|
||||
try:
|
||||
_ = compute_fvd_with_config(
|
||||
args.real_path, # noqa: F841
|
||||
args.gen_path,
|
||||
config,
|
||||
verbose=not args.quiet)
|
||||
|
||||
return 0
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error: {e}", file=sys.stderr)
|
||||
traceback.print_exc(file=sys.stderr)
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
sys.exit(main())
|
||||
@@ -1,264 +0,0 @@
|
||||
"""
|
||||
Pluggable Feature Extractors for FVD Computation.
|
||||
Supports I3D (standard), CLIP, and VideoMAE via a common interface.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from abc import ABC, abstractmethod
|
||||
from huggingface_hub import hf_hub_download
|
||||
from tqdm import tqdm
|
||||
|
||||
try:
|
||||
from transformers import CLIPModel, CLIPProcessor, VideoMAEModel
|
||||
TRANSFORMERS_AVAILABLE = True
|
||||
except ImportError:
|
||||
TRANSFORMERS_AVAILABLE = False
|
||||
|
||||
|
||||
class BaseFeatureExtractor(ABC, nn.Module):
|
||||
"""Abstract base class for all video feature extractors."""
|
||||
|
||||
def __init__(self, device: str = 'cuda'):
|
||||
super().__init__()
|
||||
self.device = torch.device(
|
||||
device if torch.cuda.is_available() else 'cpu')
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def feature_dim(self) -> int:
|
||||
"""Dimension of the output feature vector."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
videos: [B, T, C, H, W] in [0, 255] range.
|
||||
Returns:
|
||||
Preprocessed tensor ready for the model.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Extract features for a single batch.
|
||||
Args:
|
||||
videos: [B, T, C, H, W] (raw input)
|
||||
Returns:
|
||||
Features: [B, feature_dim]
|
||||
"""
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract features for a large tensor of videos by batching.
|
||||
"""
|
||||
N = len(videos)
|
||||
all_features = []
|
||||
|
||||
iterator = range(0, N, batch_size)
|
||||
if verbose:
|
||||
iterator = tqdm(
|
||||
iterator,
|
||||
desc=f"Extracting features ({self.__class__.__name__})")
|
||||
|
||||
for i in iterator:
|
||||
batch = videos[i:i + batch_size].to(self.device)
|
||||
features = self.extract_features_batch(batch)
|
||||
all_features.append(features.cpu())
|
||||
|
||||
return torch.cat(all_features, dim=0)
|
||||
|
||||
|
||||
# 1. I3D Extractor (The Standard FVD Metric)
|
||||
class I3DFeatureExtractor(BaseFeatureExtractor):
|
||||
REPO_ID = 'flateon/FVD-I3D-torchscript'
|
||||
MODEL_FILENAME = 'i3d_torchscript.pt'
|
||||
|
||||
def __init__(self, device: str = 'cuda', cache_dir: str | None = None):
|
||||
super().__init__(device)
|
||||
self.cache_dir = cache_dir
|
||||
self.model = self._load_model()
|
||||
self.model.eval()
|
||||
self.model.to(self.device)
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
return 400
|
||||
|
||||
def _load_model(self) -> torch.nn.Module:
|
||||
try:
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
return torch.jit.load(model_path, map_location=self.device)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model: {e}") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""Standard I3D preprocessing: Resize to 224, Norm to [-1, 1]."""
|
||||
B, T, C, H, W = videos.shape
|
||||
|
||||
if T < 10:
|
||||
raise ValueError(f"I3D requires at least 10 frames, got {T}")
|
||||
|
||||
# Normalize to [0, 1]
|
||||
if videos.max() > 1.0:
|
||||
videos = videos / 255.0
|
||||
|
||||
# Scale to [-1, 1]
|
||||
videos = videos * 2.0 - 1.0
|
||||
|
||||
# Resize to 224x224
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.reshape(B * T, C, H, W)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.reshape(B, T, C, 224, 224)
|
||||
|
||||
# [B, T, C, H, W] -> [B, C, T, H, W]
|
||||
return videos.permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
batch = self.preprocess(videos)
|
||||
# TorchScript I3D returns raw logits when return_features=True
|
||||
return self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
|
||||
# 2. CLIP Extractor (Semantic/Content Quality)
|
||||
class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
model_name: str = "openai/clip-vit-base-patch32"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError(
|
||||
"Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.processor = CLIPProcessor.from_pretrained(model_name)
|
||||
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
self._feature_dim = self.model.config.projection_dim
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
return self._feature_dim
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
# Ensure values are [0, 255]
|
||||
if videos.max() <= 1.0:
|
||||
videos = videos * 255.0
|
||||
|
||||
return videos.to(torch.uint8)
|
||||
|
||||
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
# Input: [B, T, C, H, W]
|
||||
B, T, C, H, W = videos.shape
|
||||
videos = self.preprocess(videos)
|
||||
|
||||
# Flatten B*T to treat frames as images
|
||||
images = videos.view(B * T, C, H, W)
|
||||
|
||||
# HF Processor
|
||||
inputs = self.processor(images=images,
|
||||
return_tensors="pt",
|
||||
padding=True)
|
||||
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
||||
|
||||
# Extract features [B*T, Dim]
|
||||
outputs = self.model.get_image_features(**inputs)
|
||||
|
||||
# Reshape [B, T, Dim] and Average Pooling over time
|
||||
outputs = outputs.view(B, T, -1)
|
||||
return outputs.mean(dim=1)
|
||||
|
||||
|
||||
# 3. VideoMAE Extractor (Structure/Motion Quality)
|
||||
class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
model_name: str = "MCG-NJU/videomae-base"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError(
|
||||
"Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
self.register_buffer(
|
||||
'mean',
|
||||
torch.tensor([0.485, 0.456, 0.406],
|
||||
device=self.device).view(1, 1, 3, 1, 1))
|
||||
self.register_buffer(
|
||||
'std',
|
||||
torch.tensor([0.229, 0.224, 0.225],
|
||||
device=self.device).view(1, 1, 3, 1, 1))
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
return self.model.config.hidden_size
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Efficient GPU-based preprocessing.
|
||||
Input: [B, T, C, H, W] in range [0, 255]
|
||||
"""
|
||||
B, T, C, H, W = videos.shape
|
||||
|
||||
# 1. Resize to 224x224
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.view(B * T, C, H, W)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.view(B, T, C, 224, 224)
|
||||
|
||||
# 2. Normalize to [0, 1]
|
||||
if videos.dtype != torch.float32:
|
||||
videos = videos.float()
|
||||
|
||||
if videos.max() > 1.0:
|
||||
videos = videos / 255.0
|
||||
|
||||
# 3. Apply ImageNet Mean/Std
|
||||
return (videos - self.mean) / self.std
|
||||
|
||||
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
# Input: [B, T, C, H, W]
|
||||
|
||||
# Fast GPU Preprocessing
|
||||
pixel_values = self.preprocess(videos)
|
||||
|
||||
# Forward pass
|
||||
outputs = self.model(pixel_values)
|
||||
|
||||
# Global Average Pooling of last hidden state [B, T_patches, 768] -> [B, 768]
|
||||
return outputs.last_hidden_state.mean(dim=1)
|
||||
|
||||
|
||||
# Factory
|
||||
def load_extractor(name: str, device: str = 'cuda') -> BaseFeatureExtractor:
|
||||
name = name.lower()
|
||||
if name == 'i3d':
|
||||
return I3DFeatureExtractor(device)
|
||||
elif name == 'clip':
|
||||
return CLIPFeatureExtractor(device)
|
||||
elif name == 'videomae':
|
||||
return VideoMAEFeatureExtractor(device)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown extractor: {name}. Options: i3d, clip, videomae")
|
||||
@@ -1,405 +0,0 @@
|
||||
import numpy as np
|
||||
import scipy.linalg
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
import pickle
|
||||
from dataclasses import dataclass, field
|
||||
from .feature_extractors import BaseFeatureExtractor, load_extractor
|
||||
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
|
||||
|
||||
|
||||
def compute_statistics(features: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Compute mean and covariance."""
|
||||
mu = np.mean(features, axis=0)
|
||||
sigma = np.cov(features, rowvar=False)
|
||||
return mu, sigma
|
||||
|
||||
|
||||
def compute_frechet_distance(mu1: np.ndarray,
|
||||
sigma1: np.ndarray,
|
||||
mu2: np.ndarray,
|
||||
sigma2: np.ndarray,
|
||||
eps: float = 1e-6) -> float:
|
||||
"""
|
||||
Compute Fréchet distance between two Gaussians.
|
||||
"""
|
||||
sigma1 = sigma1 + eps * np.eye(sigma1.shape[0])
|
||||
sigma2 = sigma2 + eps * np.eye(sigma2.shape[0])
|
||||
|
||||
diff = mu1 - mu2
|
||||
mean_distance = np.sum(diff**2)
|
||||
|
||||
trace_sum = np.trace(sigma1 + sigma2)
|
||||
|
||||
covmean = scipy.linalg.sqrtm(sigma1 @ sigma2)
|
||||
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
print(
|
||||
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
|
||||
covmean = covmean.real
|
||||
|
||||
trace_product = np.trace(covmean)
|
||||
|
||||
fvd = mean_distance + trace_sum - 2 * trace_product
|
||||
|
||||
return float(fvd)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FVDConfig:
|
||||
# default configuration for FVD computation:
|
||||
|
||||
# Video selection
|
||||
num_videos: int = 2048
|
||||
|
||||
# Feature Extractor Selection
|
||||
extractor_model: str = 'i3d' # Options: 'i3d', 'clip', 'videomae'
|
||||
|
||||
# Clip sampling
|
||||
num_frames_per_clip: int = 16
|
||||
num_clips_per_video: int = 1
|
||||
clip_strategy: str | ClipSamplingStrategy = 'beginning'
|
||||
|
||||
# Temporal subsampling
|
||||
frame_stride: int = 1 # 1=no subsampling, 2=every 2nd, 8=every 8th
|
||||
temporal_stride: int = 1 # For sliding window clips
|
||||
|
||||
# Data processing
|
||||
video_extensions: list[str] = field(
|
||||
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
|
||||
support_frame_dirs: bool = True
|
||||
|
||||
# Computation
|
||||
batch_size: int = 32
|
||||
device: str = 'cuda'
|
||||
|
||||
use_streaming: bool = True
|
||||
resize_before_extraction: bool = True
|
||||
|
||||
# Caching
|
||||
cache_real_features: str | None = None
|
||||
i3d_model_path: str | None = None
|
||||
|
||||
# Reproducibility
|
||||
seed: int | None = None
|
||||
|
||||
@classmethod
|
||||
def fvd2048_16f(cls) -> 'FVDConfig':
|
||||
"""Standard FVD protocol: 2048 videos, 16 frames, beginning clip."""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def fvd2048_128f(cls) -> 'FVDConfig':
|
||||
"""Long video protocol: 2048 videos, 128 frames."""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=128,
|
||||
clip_strategy='beginning',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def quick_test(cls) -> 'FVDConfig':
|
||||
"""Quick test config: 100 videos, 16 frames."""
|
||||
return cls(num_videos=100,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning')
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Export config to dict for logging"""
|
||||
d = self.__dict__.copy()
|
||||
d['clip_strategy'] = str(self.clip_strategy)
|
||||
return d
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""Human-readable protocol name"""
|
||||
desc = f"FVD_{self.extractor_model.upper()}_{self.num_videos}_{self.num_frames_per_clip}f"
|
||||
if self.frame_stride > 1:
|
||||
desc += f"_subsample{self.frame_stride}"
|
||||
if self.num_clips_per_video > 1:
|
||||
desc += f"_{self.num_clips_per_video}clips"
|
||||
if self.clip_strategy != 'beginning':
|
||||
desc += f"_{self.clip_strategy}"
|
||||
return desc
|
||||
|
||||
|
||||
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
|
||||
extractor: BaseFeatureExtractor,
|
||||
batch_size: int = 32,
|
||||
max_clips: int | None = None,
|
||||
verbose: bool = True) -> np.ndarray:
|
||||
"""
|
||||
Extract features from a video clip generator using streaming.
|
||||
"""
|
||||
all_features = []
|
||||
batch = []
|
||||
|
||||
if verbose:
|
||||
print(f"Extracting features with batch_size={batch_size}...")
|
||||
|
||||
with torch.no_grad():
|
||||
for clip_count, clip in enumerate(video_generator):
|
||||
batch.append(clip)
|
||||
|
||||
# Process batch when full
|
||||
if len(batch) == batch_size:
|
||||
batch_tensor = torch.stack(batch).to(extractor.device)
|
||||
features = extractor.extract_features_batch(batch_tensor)
|
||||
|
||||
all_features.append(features.detach().cpu().numpy())
|
||||
batch = []
|
||||
|
||||
if verbose and clip_count % (batch_size * 10) == 0:
|
||||
print(f"Processed {clip_count} clips...")
|
||||
|
||||
if max_clips is not None and clip_count >= max_clips:
|
||||
break
|
||||
|
||||
# Process remaining clips
|
||||
if len(batch) > 0:
|
||||
batch_tensor = torch.stack(batch).to(extractor.device)
|
||||
features = extractor.extract_features_batch(batch_tensor)
|
||||
all_features.append(features.detach().cpu().numpy())
|
||||
|
||||
if len(all_features) == 0:
|
||||
raise RuntimeError("No features extracted - check video loading")
|
||||
|
||||
features = np.concatenate(all_features, axis=0)
|
||||
|
||||
if verbose:
|
||||
print(f"Extracted {len(features)} feature vectors")
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
extractor: BaseFeatureExtractor,
|
||||
config: FVDConfig,
|
||||
cache_path: str | None = None,
|
||||
cache_name: str = "real_features") -> np.ndarray:
|
||||
"""Load features from cache or compute (with streaming support)"""
|
||||
|
||||
if cache_path is not None:
|
||||
script_dir = Path(__file__).parent
|
||||
cache_dir = script_dir / cache_path
|
||||
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
|
||||
|
||||
if cache_file.exists():
|
||||
print(f"Loading cached features from {cache_file}")
|
||||
with open(cache_file, 'rb') as f:
|
||||
features = pickle.load(f)
|
||||
|
||||
# Validate and limit based on config
|
||||
max_features = config.num_videos * config.num_clips_per_video
|
||||
|
||||
if len(features) < max_features:
|
||||
print(
|
||||
f"WARNING: Cache has {len(features)} features but need {max_features}"
|
||||
)
|
||||
print("Cached features insufficient - will recompute...")
|
||||
elif len(features) > max_features:
|
||||
print(
|
||||
f"Using {max_features} features from cache (truncated from {len(features)})"
|
||||
)
|
||||
features = features[:max_features]
|
||||
return features
|
||||
else:
|
||||
print(f"Using all {len(features)} cached features")
|
||||
return features
|
||||
|
||||
print("Computing features from scratch...")
|
||||
|
||||
if isinstance(videos, (str | Path)):
|
||||
target_size = (224, 224) if config.resize_before_extraction else None
|
||||
|
||||
video_generator = load_video_clips_streaming(
|
||||
videos,
|
||||
num_frames=config.num_frames_per_clip,
|
||||
max_videos=config.num_videos,
|
||||
clip_strategy=config.clip_strategy,
|
||||
frame_stride=config.frame_stride,
|
||||
num_clips_per_video=config.num_clips_per_video,
|
||||
video_extensions=config.video_extensions,
|
||||
support_frame_dirs=config.support_frame_dirs,
|
||||
target_size=target_size,
|
||||
verbose=True)
|
||||
|
||||
max_clips = config.num_videos * config.num_clips_per_video
|
||||
features = extract_features_streaming(video_generator,
|
||||
extractor,
|
||||
batch_size=config.batch_size,
|
||||
max_clips=max_clips,
|
||||
verbose=True)
|
||||
else:
|
||||
print(f"Extracting features from {len(videos)} video tensors...")
|
||||
features = extractor.extract_features(videos,
|
||||
batch_size=config.batch_size,
|
||||
verbose=True)
|
||||
features = features.numpy()
|
||||
|
||||
# Validate feature count
|
||||
expected_count = config.num_videos * config.num_clips_per_video
|
||||
if len(features) < expected_count:
|
||||
raise ValueError(
|
||||
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
|
||||
f"Found fewer videos than expected. Check your video directory.")
|
||||
elif len(features) > expected_count:
|
||||
print(f"Truncating {len(features)} features to {expected_count}")
|
||||
features = features[:expected_count]
|
||||
|
||||
# Cache features if requested
|
||||
if cache_path is not None:
|
||||
script_dir = Path(__file__).parent
|
||||
cache_dir = script_dir / cache_path
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
|
||||
print(f"Caching features to {cache_file}")
|
||||
with open(cache_file, 'wb') as f:
|
||||
pickle.dump(features, f)
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
|
||||
gen_videos: str | Path | torch.Tensor,
|
||||
config: FVDConfig,
|
||||
verbose: bool = True) -> dict:
|
||||
"""
|
||||
Compute FVD using a standardized configuration.
|
||||
|
||||
This is the recommended way to compute FVD for reproducibility.
|
||||
|
||||
Args:
|
||||
real_videos: Path or tensors
|
||||
gen_videos: Path or tensors
|
||||
config: FVDConfig specifying protocol
|
||||
verbose: Print progress
|
||||
|
||||
Returns:
|
||||
results: Dictionary with:
|
||||
- 'fvd': FVD score (float)
|
||||
- 'protocol': Protocol name (str)
|
||||
- 'model': Feature extractor model name (str)
|
||||
- 'config': Configuration dict
|
||||
|
||||
Example:
|
||||
>>> config = FVDConfig.fvd2048_16f()
|
||||
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
>>> print(f"FVD: {results['fvd']:.2f}")
|
||||
"""
|
||||
|
||||
# Seed for reproducibility
|
||||
if config.seed is not None:
|
||||
import random as _rnd
|
||||
_rnd.seed(config.seed)
|
||||
np.random.seed(config.seed)
|
||||
torch.manual_seed(config.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(config.seed)
|
||||
|
||||
if verbose:
|
||||
print("=" * 70)
|
||||
print(f"Computing FVD with protocol: {config}")
|
||||
print(f"Model: {config.extractor_model.upper()}")
|
||||
print("=" * 70)
|
||||
print("\nConfiguration:")
|
||||
for key, value in config.to_dict().items():
|
||||
print(f" {key}: {value}")
|
||||
print()
|
||||
|
||||
# Initialize Extractor using Factory
|
||||
if verbose:
|
||||
print(
|
||||
f"\nInitializing {config.extractor_model.upper()} model on {config.device}..."
|
||||
)
|
||||
|
||||
extractor = load_extractor(config.extractor_model, device=config.device)
|
||||
|
||||
# Extract features
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting REAL video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
real_features = load_or_compute_features(
|
||||
videos=real_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=config.cache_real_features,
|
||||
cache_name="real_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting GENERATED video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
gen_features = load_or_compute_features(videos=gen_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=None,
|
||||
cache_name="gen_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\nReal videos/clips: {len(real_features)}")
|
||||
print(f"Generated videos/clips: {len(gen_features)}")
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing statistics...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
mu_real, sigma_real = compute_statistics(real_features)
|
||||
mu_gen, sigma_gen = compute_statistics(gen_features)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing Fréchet distance...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
fvd = compute_frechet_distance(mu_real, sigma_real, mu_gen, sigma_gen)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print(f"FVD Score ({config.extractor_model.upper()}): {fvd:.4f}")
|
||||
print(f"Protocol: {config}")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
results = {
|
||||
'fvd': fvd,
|
||||
'protocol': str(config),
|
||||
'model': config.extractor_model,
|
||||
'config': config.to_dict(),
|
||||
}
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def compute_fvd(real_videos: str | Path | torch.Tensor,
|
||||
gen_videos: str | Path | torch.Tensor,
|
||||
num_frames: int = 16,
|
||||
batch_size: int = 32,
|
||||
device: str = 'cuda',
|
||||
num_videos: int | None = 2048,
|
||||
cache_real_features: str | None = None,
|
||||
i3d_model_path: str | None = None,
|
||||
seed: int | None = None,
|
||||
verbose: bool = True) -> float:
|
||||
"""
|
||||
Backward compatibility wrapper for computing FVD (defaults to I3D).
|
||||
"""
|
||||
num_videos = num_videos if num_videos is not None else 2048
|
||||
|
||||
config = FVDConfig(
|
||||
num_videos=num_videos,
|
||||
num_frames_per_clip=num_frames,
|
||||
extractor_model='i3d', # Default to I3D
|
||||
batch_size=batch_size,
|
||||
device=device,
|
||||
cache_real_features=cache_real_features,
|
||||
i3d_model_path=i3d_model_path,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
|
||||
return result['fvd']
|
||||
@@ -1,142 +0,0 @@
|
||||
"""I3D Feature Extractor for FVD Computation"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pathlib import Path
|
||||
from huggingface_hub import hf_hub_download
|
||||
from tqdm import tqdm
|
||||
from contextlib import suppress
|
||||
|
||||
|
||||
class I3DFeatureExtractor(nn.Module):
|
||||
"""
|
||||
I3D feature extractor for FVD computation.
|
||||
Extracts 400-dimensional features from videos using I3D model
|
||||
trained on Kinetics-400.
|
||||
"""
|
||||
|
||||
REPO_ID = 'flateon/FVD-I3D-torchscript'
|
||||
MODEL_FILENAME = 'i3d_torchscript.pt'
|
||||
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
cache_dir: str | Path | None = None):
|
||||
super().__init__()
|
||||
|
||||
self.device_str = device
|
||||
if device == 'cuda' and not torch.cuda.is_available():
|
||||
print(
|
||||
"Warning: CUDA requested but not available – falling back to CPU"
|
||||
)
|
||||
self.device = torch.device('cpu')
|
||||
else:
|
||||
self.device = torch.device(device)
|
||||
|
||||
self.cache_dir: str | None
|
||||
if cache_dir is not None:
|
||||
self.cache_dir = str(Path(cache_dir).resolve())
|
||||
else:
|
||||
self.cache_dir = None # Use HF default cache
|
||||
|
||||
self.model = self._load_model()
|
||||
self.model.eval()
|
||||
|
||||
with suppress(Exception):
|
||||
self.model.to(self.device)
|
||||
|
||||
def _load_model(self) -> torch.nn.Module:
|
||||
"""Download and load I3D TorchScript model from Hugging Face Hub."""
|
||||
print(f"Loading I3D model from Hugging Face Hub ({self.REPO_ID})...")
|
||||
|
||||
try:
|
||||
# Download model from Hugging Face Hub
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
|
||||
# Load directly to chosen device
|
||||
model = torch.jit.load(model_path, map_location=self.device)
|
||||
print("I3D model loaded successfully")
|
||||
return model
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Preprocess videos for I3D.
|
||||
|
||||
Args:
|
||||
videos: [B, T, C, H, W], values in [0, 255]
|
||||
|
||||
Returns:
|
||||
Preprocessed videos [B, C, T, 224, 224] (normalized and resized)
|
||||
"""
|
||||
B, T, C, H, W = videos.shape
|
||||
|
||||
if T < 10:
|
||||
raise ValueError(f"I3D requires at least 10 frames, got {T}")
|
||||
|
||||
# Normalize to [0, 1] if needed
|
||||
if videos.max() > 1.0:
|
||||
videos = videos / 255.0
|
||||
|
||||
# Resize to 224x224 if needed
|
||||
if H != 224 or W != 224:
|
||||
videos = videos.reshape(B * T, C, H, W)
|
||||
videos = F.interpolate(videos,
|
||||
size=(224, 224),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
videos = videos.reshape(B, T, C, 224, 224)
|
||||
|
||||
# Convert to [B, C, T, H, W] format
|
||||
videos = videos.permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
return videos
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract I3D features
|
||||
|
||||
Args:
|
||||
videos: [N, T, C, H, W], values in [0, 255]
|
||||
batch_size: Batch size for processing
|
||||
verbose: Show progress bar
|
||||
|
||||
Returns:
|
||||
Features [N, 400]
|
||||
"""
|
||||
N = len(videos)
|
||||
all_features = []
|
||||
|
||||
iterator = range(0, N, batch_size)
|
||||
if verbose:
|
||||
iterator = tqdm(iterator, desc="Extracting I3D features")
|
||||
|
||||
for i in iterator:
|
||||
batch = videos[i:i + batch_size].to(self.device)
|
||||
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
|
||||
|
||||
# Use the HF model without rescale/resize (we handle it in preprocess)
|
||||
features = self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
all_features.append(features.cpu())
|
||||
|
||||
return torch.cat(all_features, dim=0)
|
||||
|
||||
def __call__(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32) -> torch.Tensor:
|
||||
return self.extract_features(videos, batch_size=batch_size)
|
||||
@@ -1,54 +0,0 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
root_dir = Path(__file__).parent.parent.parent
|
||||
sys.path.insert(0, str(root_dir))
|
||||
|
||||
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config # noqa: E402
|
||||
|
||||
|
||||
def main() -> None:
|
||||
script_dir = Path(__file__).parent.resolve()
|
||||
|
||||
# Define directories
|
||||
real_dir = "benchmarks/data/real_videos"
|
||||
gen_dir = "benchmarks/data/generated_videos"
|
||||
|
||||
# Compare all 3 models
|
||||
models_to_test = ['i3d', 'clip', 'videomae']
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print("STARTING COMPARISON BENCHMARK")
|
||||
print(f"{'='*60}")
|
||||
|
||||
for model_name in models_to_test:
|
||||
print(f"\n>>> Running evaluation with {model_name.upper()}...")
|
||||
|
||||
try:
|
||||
cfg = FVDConfig(
|
||||
num_videos=650,
|
||||
num_frames_per_clip=16,
|
||||
extractor_model=model_name,
|
||||
clip_strategy='beginning',
|
||||
device='cuda',
|
||||
seed=42,
|
||||
# Use separate cache folders for each model to avoid conflicts
|
||||
cache_real_features=str(script_dir / f'fvd-cache/{model_name}'),
|
||||
)
|
||||
|
||||
results = compute_fvd_with_config(real_dir,
|
||||
gen_dir,
|
||||
cfg,
|
||||
verbose=False)
|
||||
print(f"FVD: {results['fvd']}\nModel: {results['model']}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"{model_name.upper()} Failed: {e}")
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print("BENCHMARK COMPLETE")
|
||||
print(f"{'='*60}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,97 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import random
|
||||
from fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
script_path = Path(__file__).resolve()
|
||||
fastvideo_root = script_path.parent.parent.parent
|
||||
sys.path.insert(0, str(fastvideo_root))
|
||||
|
||||
|
||||
def split_videos(video_dir: Path, n_per_subset: int = 128, seed: int = 42):
|
||||
subset_a = video_dir.parent / 'bair_full_subset_A'
|
||||
subset_b = video_dir.parent / 'bair_full_subset_B'
|
||||
|
||||
if subset_a.exists():
|
||||
shutil.rmtree(subset_a)
|
||||
if subset_b.exists():
|
||||
shutil.rmtree(subset_b)
|
||||
|
||||
subset_a.mkdir(parents=True)
|
||||
subset_b.mkdir(parents=True)
|
||||
|
||||
videos = sorted(video_dir.glob('*.mp4'))
|
||||
|
||||
random.seed(seed)
|
||||
shuffled = list(videos)
|
||||
random.shuffle(shuffled)
|
||||
|
||||
needed = n_per_subset * 2
|
||||
if len(shuffled) > needed:
|
||||
shuffled = shuffled[:needed]
|
||||
|
||||
mid = len(shuffled) // 2
|
||||
|
||||
print(f"\nSplitting {len(shuffled)} BAIR FULL videos:")
|
||||
print(f" Subset A: {mid} videos")
|
||||
print(f" Subset B: {len(shuffled) - mid} videos")
|
||||
|
||||
for v in shuffled[:mid]:
|
||||
shutil.copy2(v, subset_a / v.name)
|
||||
|
||||
for v in shuffled[mid:]:
|
||||
shutil.copy2(v, subset_b / v.name)
|
||||
|
||||
return subset_a, subset_b, mid
|
||||
|
||||
|
||||
def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
config = FVDConfig(num_videos=num_videos,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
batch_size=8,
|
||||
device='cuda',
|
||||
seed=42)
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 1: Identity Test")
|
||||
print("=" * 70)
|
||||
|
||||
result1 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_a),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_identity = result1['fvd']
|
||||
print(f"\nIdentity FVD: {fvd_identity:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 2: Real vs Real")
|
||||
print("=" * 70)
|
||||
|
||||
result2 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_b),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_real = result2['fvd']
|
||||
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("RESULTS")
|
||||
print("=" * 70)
|
||||
print(f"Identity: {fvd_identity:.2f}")
|
||||
print(f"Real vs Real: {fvd_real:.2f}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
bair_dir = Path('benchmarks/data/bair_full_videos')
|
||||
|
||||
subset_a, subset_b, count = split_videos(bair_dir,
|
||||
n_per_subset=128,
|
||||
seed=42)
|
||||
validate_fvd(subset_a, subset_b, count)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,490 +0,0 @@
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
from tqdm import tqdm
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ClipSamplingStrategy(Enum):
|
||||
"""Clip sampling strategies for FVD evaluation."""
|
||||
BEGINNING = 'beginning' # Take first N frames (most common)
|
||||
RANDOM = 'random' # Random N consecutive frames
|
||||
UNIFORM = 'uniform' # Uniformly spaced frames across video
|
||||
MIDDLE = 'middle' # Middle N frames
|
||||
SLIDING = 'sliding' # Multiple sliding windows
|
||||
ALL = 'all' # All possible clips
|
||||
|
||||
|
||||
def _load_video_cv2(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Load video from video file using OpenCV.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file (MP4, AVI, MOV, MKV)
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
video_path = str(video_path)
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
|
||||
frames = []
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
if num_frames is None:
|
||||
# Read all available frames
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
if len(frames) == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
return frames
|
||||
|
||||
if total_frames == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
# Determine frame indices for sampling
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(
|
||||
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Extract frames
|
||||
for idx in frame_indices:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
||||
ret, frame = cap.read()
|
||||
|
||||
if not ret:
|
||||
if len(frames) > 0:
|
||||
frames.append(frames[-1].copy())
|
||||
else:
|
||||
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
frames.append(np.zeros((h, w, 3), dtype=np.uint8))
|
||||
continue
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _load_video_from_frames(
|
||||
frame_dir: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform',
|
||||
frame_extensions: list[str] | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Load video from directory of frame images.
|
||||
|
||||
Args:
|
||||
frame_dir: Directory containing frames
|
||||
num_frames: Number of frames to sample
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
frame_extensions: Image file extensions to look for
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
if frame_extensions is None:
|
||||
frame_extensions = ['.jpg', '.png', '.jpeg', '.bmp']
|
||||
|
||||
frame_dir = Path(frame_dir)
|
||||
|
||||
if not frame_dir.exists():
|
||||
raise FileNotFoundError(f"Frame directory not found: {frame_dir}")
|
||||
|
||||
# Find all frames
|
||||
frame_files: list[Path] = []
|
||||
for ext in frame_extensions:
|
||||
frame_files.extend(frame_dir.glob(f"*{ext}"))
|
||||
|
||||
if len(frame_files) == 0:
|
||||
raise ValueError(
|
||||
f"No frames found in {frame_dir} with extensions {frame_extensions}"
|
||||
)
|
||||
|
||||
frame_files = sorted(frame_files, key=lambda x: x.name)
|
||||
total_frames = len(frame_files)
|
||||
|
||||
# Determine frame indices
|
||||
if num_frames is None:
|
||||
frame_indices = list(range(total_frames))
|
||||
else:
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
|
||||
num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0,
|
||||
total_frames - 1,
|
||||
num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Load frames
|
||||
frames = []
|
||||
for idx in frame_indices:
|
||||
frame_path = frame_files[idx]
|
||||
frame = cv2.imread(str(frame_path))
|
||||
|
||||
if frame is None:
|
||||
raise RuntimeError(f"Failed to load frame: {frame_path}")
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
# Stack and convert to tensor
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _detect_video_format(path: str | Path) -> str:
|
||||
"""
|
||||
Detect if path is a video file or frame directory.
|
||||
|
||||
Returns:
|
||||
'video_file', 'frame_directory', or 'unknown'
|
||||
"""
|
||||
path = Path(path)
|
||||
|
||||
if path.is_file():
|
||||
return 'video_file'
|
||||
elif path.is_dir():
|
||||
# Check if contains image files
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(path.glob(f"*{ext}")):
|
||||
return 'frame_directory'
|
||||
return 'unknown'
|
||||
else:
|
||||
raise ValueError(f"Path does not exist: {path}")
|
||||
|
||||
|
||||
def load_video_auto(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Automatically detect format and load video.
|
||||
|
||||
Supports:
|
||||
- Video files (MP4, AVI, MOV, MKV)
|
||||
- Frame directories (JPG, PNG)
|
||||
|
||||
Args:
|
||||
video_path: Path to video file or frame directory
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
format_type = _detect_video_format(video_path)
|
||||
|
||||
if format_type == 'video_file':
|
||||
return _load_video_cv2(video_path, num_frames, sample_strategy)
|
||||
elif format_type == 'frame_directory':
|
||||
return _load_video_from_frames(video_path, num_frames, sample_strategy)
|
||||
else:
|
||||
raise ValueError(f"Unknown video format at {video_path}")
|
||||
|
||||
|
||||
def sample_clips_from_video(
|
||||
video: torch.Tensor,
|
||||
num_frames_per_clip: int = 16,
|
||||
num_clips: int = 1,
|
||||
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
|
||||
frame_stride: int = 1,
|
||||
temporal_stride: int = 1) -> list[torch.Tensor]:
|
||||
"""
|
||||
Sample clips from a video with various strategies.
|
||||
|
||||
Args:
|
||||
video: [T, C, H, W] full video
|
||||
num_frames_per_clip: Frames per clip
|
||||
num_clips: Number of clips to extract
|
||||
strategy: ClipSamplingStrategy or string ('beginning', 'random', etc.)
|
||||
frame_stride: Skip frames (FPS control: 1=all, 2=every 2nd, 8=every 8th)
|
||||
temporal_stride: Stride between clips for sliding window
|
||||
|
||||
Returns:
|
||||
List of clips, each [num_frames_per_clip, C, H, W]
|
||||
|
||||
Examples:
|
||||
>>> # Beginning clip (most common for FVD)
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='beginning')
|
||||
|
||||
>>> # Multiple random clips
|
||||
>>> clips = sample_clips_from_video(video, 16, num_clips=4, strategy='random')
|
||||
|
||||
>>> # Subsample FPS by 2x (every 2nd frame)
|
||||
>>> clips = sample_clips_from_video(video, 16, frame_stride=2)
|
||||
|
||||
>>> # Sliding window with overlap
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='sliding', temporal_stride=8)
|
||||
"""
|
||||
# Convert string to enum if needed
|
||||
if isinstance(strategy, str):
|
||||
strategy = ClipSamplingStrategy(strategy)
|
||||
|
||||
T, C, H, W = video.shape
|
||||
|
||||
# Apply frame stride (FPS subsampling)
|
||||
if frame_stride > 1:
|
||||
video = video[::frame_stride]
|
||||
T = len(video)
|
||||
|
||||
effective_clip_length = num_frames_per_clip
|
||||
|
||||
# Handle videos shorter than clip length
|
||||
if effective_clip_length > T:
|
||||
pad_length = effective_clip_length - T
|
||||
last_frame = video[-1:].repeat(pad_length, 1, 1, 1)
|
||||
video = torch.cat([video, last_frame], dim=0)
|
||||
T = len(video)
|
||||
|
||||
clips = []
|
||||
|
||||
if strategy == ClipSamplingStrategy.BEGINNING:
|
||||
# Take first clip (most common for FVD evaluation)
|
||||
clip = video[:effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.MIDDLE:
|
||||
# Take middle clip
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.RANDOM:
|
||||
# Sample N random clips
|
||||
for _ in range(num_clips):
|
||||
if effective_clip_length == T:
|
||||
start = 0
|
||||
else:
|
||||
start = np.random.randint(0, T - effective_clip_length + 1)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.UNIFORM:
|
||||
# Uniformly spaced clips
|
||||
if num_clips == 1:
|
||||
# Single clip from middle
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
else:
|
||||
# Multiple uniformly spaced clips
|
||||
step = (T - effective_clip_length) / (num_clips -
|
||||
1) if num_clips > 1 else 0
|
||||
for i in range(num_clips):
|
||||
start = int(i * step)
|
||||
start = min(start, T - effective_clip_length)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.SLIDING:
|
||||
# Sliding window with stride
|
||||
for start in range(0, T - effective_clip_length + 1, temporal_stride):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
if len(clips) >= num_clips:
|
||||
break
|
||||
|
||||
elif strategy == ClipSamplingStrategy.ALL:
|
||||
# All possible clips (overlapping)
|
||||
for start in range(T - effective_clip_length + 1):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {strategy}")
|
||||
|
||||
return clips
|
||||
|
||||
|
||||
def load_video_clips_streaming(directory: str | Path,
|
||||
num_frames: int = 16,
|
||||
max_videos: int | None = None,
|
||||
clip_strategy: str
|
||||
| ClipSamplingStrategy = 'beginning',
|
||||
frame_stride: int = 1,
|
||||
num_clips_per_video: int = 1,
|
||||
video_extensions: list[str] | None = None,
|
||||
support_frame_dirs: bool = True,
|
||||
target_size: tuple[int, int] | None = (224, 224),
|
||||
verbose: bool = True) -> Iterator[torch.Tensor]:
|
||||
"""
|
||||
This generator yields clips one-by-one instead of loading all videos into RAM.
|
||||
Perfect for large datasets where memory is limited.
|
||||
|
||||
Args:
|
||||
directory: Path to directory with videos
|
||||
num_frames: Frames per clip
|
||||
max_videos: Max videos to load
|
||||
clip_strategy: 'beginning', 'random', 'uniform', etc.
|
||||
frame_stride: Frame skip (1=all, 2=every 2nd, 8=every 8th)
|
||||
num_clips_per_video: Number of clips per video
|
||||
video_extensions: Video file extensions
|
||||
support_frame_dirs: Also load frame directories
|
||||
target_size: Resize clips to (H, W). If None, keep original size.
|
||||
verbose: Show progress
|
||||
|
||||
Yields:
|
||||
clip: [T, C, H, W] individual clips
|
||||
|
||||
Example:
|
||||
>>> for clip in load_video_clips_streaming('data/videos/', num_frames=16):
|
||||
>>> features = model.extract_features(clip.unsqueeze(0))
|
||||
>>> # Process one clip at a time - low memory usage!
|
||||
"""
|
||||
if video_extensions is None:
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.mkv']
|
||||
|
||||
directory = Path(directory)
|
||||
|
||||
if not directory.exists():
|
||||
raise FileNotFoundError(f"Directory not found: {directory}")
|
||||
|
||||
# Find video paths
|
||||
video_paths: list[Path] = []
|
||||
|
||||
# Find video files
|
||||
for ext in video_extensions:
|
||||
video_paths.extend(directory.glob(f"**/*{ext}"))
|
||||
|
||||
# Find frame directories if enabled
|
||||
if support_frame_dirs:
|
||||
for subdir in directory.iterdir():
|
||||
if subdir.is_dir():
|
||||
# Check if it contains frames
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(subdir.glob(f"*{ext}")):
|
||||
video_paths.append(subdir)
|
||||
break
|
||||
|
||||
if len(video_paths) == 0:
|
||||
raise ValueError(f"No videos found in {directory}")
|
||||
|
||||
video_paths = sorted(video_paths)
|
||||
|
||||
if max_videos is not None:
|
||||
video_paths = video_paths[:max_videos]
|
||||
|
||||
if verbose:
|
||||
print(f"Found {len(video_paths)} videos in {directory}")
|
||||
if num_clips_per_video > 1:
|
||||
print(f"Extracting {num_clips_per_video} clips per video...")
|
||||
if frame_stride > 1:
|
||||
print(f"Subsampling frames with stride {frame_stride}...")
|
||||
if target_size:
|
||||
print(f"Resizing clips to {target_size}...")
|
||||
|
||||
# Track statistics
|
||||
failed_count = 0
|
||||
total_clips = 0
|
||||
|
||||
iterator = tqdm(video_paths,
|
||||
desc="Loading videos") if verbose else video_paths
|
||||
|
||||
for video_path in iterator:
|
||||
try:
|
||||
# Load full video
|
||||
video = load_video_auto(video_path,
|
||||
num_frames=None,
|
||||
sample_strategy='uniform')
|
||||
|
||||
# Sample clips from video
|
||||
clips = sample_clips_from_video(video,
|
||||
num_frames_per_clip=num_frames,
|
||||
num_clips=num_clips_per_video,
|
||||
strategy=clip_strategy,
|
||||
frame_stride=frame_stride)
|
||||
|
||||
if target_size is not None:
|
||||
resized_clips = []
|
||||
for clip in clips:
|
||||
T, C, H, W = clip.shape
|
||||
if target_size != (H, W):
|
||||
# Resize to target size
|
||||
clip = clip.contiguous(
|
||||
) # Fix non-contiguous tensors first
|
||||
clip_flat = clip.view(T * C, H,
|
||||
W).unsqueeze(0) # [1, T*C, H, W]
|
||||
clip_resized = torch.nn.functional.interpolate(
|
||||
clip_flat,
|
||||
size=target_size,
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
clip = clip_resized.squeeze(0).view(
|
||||
T, C, target_size[0],
|
||||
target_size[1]) # Back to [T, C, H, W]
|
||||
resized_clips.append(clip)
|
||||
clips = resized_clips
|
||||
|
||||
# Yield clips one by one
|
||||
for clip in clips:
|
||||
yield clip
|
||||
total_clips += 1
|
||||
|
||||
# Free memory
|
||||
del video, clips
|
||||
|
||||
except Exception as e:
|
||||
failed_count += 1
|
||||
if verbose:
|
||||
print(f"\nWarning: Failed to load {video_path}: {e}")
|
||||
continue
|
||||
|
||||
# Validate
|
||||
if total_clips == 0:
|
||||
raise RuntimeError(f"Failed to load any videos from {directory}")
|
||||
|
||||
failure_rate = failed_count / len(video_paths)
|
||||
if failure_rate > 0.1: # More than 10% failed
|
||||
print(
|
||||
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
|
||||
)
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
|
||||
)
|
||||
@@ -1,7 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless transformers huggingface_hub
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
@@ -1,4 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
+5
-5
@@ -9,12 +9,11 @@ import datetime
|
||||
import locale
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
# 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+).
|
||||
# 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`
|
||||
from collections import namedtuple
|
||||
|
||||
from fastvideo.envs import environment_variables
|
||||
@@ -496,7 +495,8 @@ def get_pip_packages(run_lambda, patterns=None):
|
||||
|
||||
if pip_available:
|
||||
cmd = [sys.executable, '-mpip', 'list', '--format=freeze']
|
||||
elif shutil.which("uv") is not None:
|
||||
elif os.environ.get("UV") is not None:
|
||||
print("uv is set")
|
||||
cmd = ["uv", "pip", "list", "--format=freeze"]
|
||||
else:
|
||||
raise RuntimeError(
|
||||
|
||||
@@ -443,6 +443,7 @@
|
||||
1025,
|
||||
"fixed",
|
||||
24,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
@@ -490,6 +491,11 @@
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "X://insert/path/here.mp4"
|
||||
},
|
||||
"enable_teacache": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -636,4 +642,4 @@
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
}
|
||||
@@ -353,7 +353,8 @@
|
||||
1024,
|
||||
"fixed",
|
||||
24,
|
||||
"X://insert/path/here.mp4"
|
||||
"X://insert/path/here.mp4",
|
||||
true
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"height": {
|
||||
@@ -400,6 +401,11 @@
|
||||
"isAuto": true,
|
||||
"value": "X://insert/path/here.mp4",
|
||||
"cachedValue": "X://insert/path/here.mp4"
|
||||
},
|
||||
"enable_teacache": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -688,4 +694,4 @@
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
}
|
||||
@@ -31,6 +31,9 @@ class InferenceArgs:
|
||||
"image_path": ("STRING", {
|
||||
"default": "X://insert/path/here.mp4"
|
||||
}),
|
||||
"enable_teacache": ([True, False], {
|
||||
"default": False
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,6 +57,7 @@ class InferenceArgs:
|
||||
seed,
|
||||
fps,
|
||||
image_path,
|
||||
enable_teacache,
|
||||
):
|
||||
raw_args = {
|
||||
"height": height,
|
||||
@@ -65,6 +69,7 @@ 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
|
||||
|
||||
@@ -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"
|
||||
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
|
||||
]
|
||||
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
|
||||
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -0,0 +1,145 @@
|
||||
import os
|
||||
from collections import defaultdict
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from st_attn import sliding_tile_attention
|
||||
from triton.testing import do_bench
|
||||
|
||||
|
||||
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
assert mode in ["fwd", "bwd", "fwd_bwd"]
|
||||
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
|
||||
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
|
||||
|
||||
|
||||
def compute_TFLOPS(flops, ms):
|
||||
flops = flops / 1e12
|
||||
ms = ms / 1e3
|
||||
return flops / ms
|
||||
|
||||
|
||||
def benchmark_attention(configurations):
|
||||
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
|
||||
|
||||
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
|
||||
print("=" * 60)
|
||||
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
|
||||
|
||||
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
|
||||
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
|
||||
# # Warmup for forward pass
|
||||
# for _ in range(10):
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
|
||||
# # Time the forward pass
|
||||
# for i in range(10):
|
||||
# start_events_fwd[i].record()
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
# end_events_fwd[i].record()
|
||||
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
|
||||
|
||||
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
# time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
|
||||
results['fwd'][(D, causal)].append((N, tflops_fwd))
|
||||
|
||||
print(f"Average time for forward pass (ms): {ms:.2f}")
|
||||
print(f"Average TFLOPS: {tflops_fwd}")
|
||||
print("-" * 60)
|
||||
|
||||
# torch.cuda.empty_cache()
|
||||
# torch.cuda.synchronize()
|
||||
|
||||
# # Prepare for timing backward pass
|
||||
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
|
||||
# # Warmup for backward pass
|
||||
# for _ in range(10):
|
||||
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
|
||||
|
||||
# # Time the backward pass
|
||||
# for i in range(10):
|
||||
# start_events_bwd[i].record()
|
||||
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
|
||||
# end_events_bwd[i].record()
|
||||
|
||||
# torch.cuda.synchronize()
|
||||
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
|
||||
# time_us_bwd = np.mean(times_bwd) * 1000
|
||||
|
||||
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
|
||||
# results['bwd'][(D, causal)].append((N, tflops_bwd))
|
||||
|
||||
# print(f"Average time for backward pass(ms): {ms:.2f}")
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def plot_results(results):
|
||||
os.makedirs('benchmark_results', exist_ok=True)
|
||||
for mode in ['fwd', 'bwd']:
|
||||
for (D, causal), values in results[mode].items():
|
||||
seq_lens = [x[0] for x in values]
|
||||
tflops = [x[1] for x in values]
|
||||
|
||||
plt.figure(figsize=(10, 6))
|
||||
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
|
||||
plt.xlabel('Sequence Length')
|
||||
plt.ylabel('TFLOPS')
|
||||
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
|
||||
plt.grid(True)
|
||||
|
||||
# Adding the numerical y value on top of each bar
|
||||
for bar in bars:
|
||||
yval = bar.get_height()
|
||||
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
|
||||
|
||||
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
|
||||
plt.savefig(filename)
|
||||
plt.close()
|
||||
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
|
||||
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
# (16, 16, 768*8, 128, False),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768, 128, True),
|
||||
# (16, 16, 768*2, 128, True),
|
||||
# (16, 16, 768*4, 128, True),
|
||||
# (16, 16, 768*8, 128, True),
|
||||
# (16, 16, 768*16, 128, True),
|
||||
# (16, 32, 768, 64, False),
|
||||
# (16, 32, 768*2, 64, False),
|
||||
# (16, 32, 768*4, 64, False),
|
||||
# (16, 32, 768*8, 64, False),
|
||||
# (16, 32, 768*16, 64, False),
|
||||
# (16, 32, 768, 64, True),
|
||||
# (16, 32, 768*2, 64, True),
|
||||
# (16, 32, 768*4, 64, True),
|
||||
# (16, 32, 768*8, 64, True),
|
||||
# (16, 32, 768*16, 64, True),
|
||||
]
|
||||
|
||||
results = benchmark_attention(configurations)
|
||||
# plot_results(results)
|
||||
@@ -0,0 +1,224 @@
|
||||
import torch
|
||||
import argparse
|
||||
from triton.testing import do_bench
|
||||
from vsa import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
import triton
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward and backward passes."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
|
||||
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
fwd_time = do_bench(
|
||||
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
# Backward pass
|
||||
grad_output = torch.randn_like(o)
|
||||
|
||||
# Warm-up runs
|
||||
for _ in range(5):
|
||||
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
bwd_time = do_bench(
|
||||
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
|
||||
|
||||
return sparse_tflops, sparse_bwd_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,217 @@
|
||||
import torch
|
||||
import argparse
|
||||
import triton.testing
|
||||
from vsa import block_sparse_attn
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward+backward pass."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION FORWARD+BACKWARD BENCHMARK ===")
|
||||
|
||||
# Combined forward+backward pass
|
||||
# Warm-up run
|
||||
q_fwd = q.clone().requires_grad_(True)
|
||||
k_fwd = k.clone().requires_grad_(True)
|
||||
v_fwd = v.clone().requires_grad_(True)
|
||||
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
grad_output = torch.randn_like(o)
|
||||
o.backward(grad_output)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward+backward
|
||||
def forward_backward_fn():
|
||||
q_fwd = q.clone().requires_grad_(True)
|
||||
k_fwd = k.clone().requires_grad_(True)
|
||||
v_fwd = v.clone().requires_grad_(True)
|
||||
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
grad_output = torch.randn_like(o)
|
||||
o.backward(grad_output)
|
||||
|
||||
total_time = triton.testing.do_bench(
|
||||
forward_backward_fn,
|
||||
warmup=25,
|
||||
rep=100,
|
||||
return_mode='mean'
|
||||
)
|
||||
|
||||
# Total flops for forward + backward (forward + 2.5x backward approximation)
|
||||
total_flops = flops + 2.5 * flops # 3.5x the forward flops
|
||||
sparse_tflops = total_flops / total_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
return sparse_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_fwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_sta.py
|
||||
@@ -0,0 +1,96 @@
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python ../benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
|
||||
## STA Configuration Logic
|
||||
Here is a diagram of how the window is configured and passed through the FastVideo pipeline:
|
||||
|
||||
<div align="center">
|
||||
<img src="../../../docs/assets/images/STA_configuration.png" width="80%"/>
|
||||
</div>
|
||||
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -0,0 +1,15 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'st_attn': {
|
||||
'source_files': {
|
||||
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['st_attn']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -0,0 +1,76 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.6"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('st_attn root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['st_attn.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -0,0 +1,23 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
try:
|
||||
from st_attn_cuda import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
seq_length = q_all.shape[2]
|
||||
dit_seq_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
assert q_all.shape[2] == 82944
|
||||
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
|
||||
assert q_all.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
|
||||
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_shape]
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
for batch in range(q_all.shape[0]):
|
||||
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
|
||||
k_all[batch:batch + 1,
|
||||
head_index:head_index + 1], v_all[batch:batch + 1,
|
||||
head_index:head_index + 1],
|
||||
hidden_states[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
+347
-79
@@ -451,10 +451,6 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
// Shared memory size for the kernel.
|
||||
// We use the maximum available shared memory (kittens::MAX_SHARED_MEMORY)
|
||||
// which is approximately 227KB on H100, necessary for the high-performance
|
||||
// TMA-based attention tiles with multiple stages.
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
@@ -462,31 +458,104 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
dim3 grid_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(2, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 3, 0); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(2, 2, 3); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 0, 5); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true,1, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==5 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==5 && kernel_h_size == 3 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true,2, 0, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
@@ -499,67 +568,266 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(3, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 3, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 3, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 0, 0); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 0, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 3, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 0, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(0, 3, 0); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
else if (kernel_aspect_ratio_flag == 3) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 3, 0); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(1, 2, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(1, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 2, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 2, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(0, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 0, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(0, 3, 2); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
|
||||
else {
|
||||
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
|
||||
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch
|
||||
from .support_flex_sta import get_sliding_tile_attention_mask
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from flex_sta_ref import get_sliding_tile_attention_mask
|
||||
from st_attn import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
# from flash_attn_interface import flash_attn_func
|
||||
from tqdm import tqdm
|
||||
@@ -73,23 +73,15 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
|
||||
|
||||
# Example usage
|
||||
def test_sliding_tile_attention():
|
||||
if not torch.cuda.is_available():
|
||||
return
|
||||
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
|
||||
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_sliding_tile_attention()
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
|
||||
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
@@ -0,0 +1,156 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
# Add the parent directory to the path to import block_sparse_attn
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from tests.utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
|
||||
from vsa import block_sparse_attn
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
q_ = Q.clone().float().requires_grad_()
|
||||
k_ = K.clone().float().requires_grad_()
|
||||
v_ = V.clone().float().requires_grad_()
|
||||
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
dO_ = dO
|
||||
output.backward(dO_)
|
||||
return (
|
||||
output.to(torch.bfloat16),
|
||||
q_.grad.to(torch.bfloat16),
|
||||
k_.grad.to(torch.bfloat16),
|
||||
v_.grad.to(torch.bfloat16),
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
|
||||
Q = Q.detach().requires_grad_()
|
||||
K = K.detach().requires_grad_()
|
||||
V = V.detach().requires_grad_()
|
||||
|
||||
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
|
||||
output = output[:, :, non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def get_non_pad_index(
|
||||
vid_len: torch.LongTensor,
|
||||
n_win: int,
|
||||
win_size: int,
|
||||
):
|
||||
device = vid_len.device
|
||||
starts_pad = torch.arange(n_win, device=device) * win_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
|
||||
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
|
||||
|
||||
return index_pad[index_mask]
|
||||
|
||||
def generate_tensor(shape, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
return tensor
|
||||
|
||||
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
|
||||
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
|
||||
|
||||
|
||||
def vsa_pad(x, non_pad_index, num_blocks, block_size):
|
||||
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
|
||||
padded_x[:, :, non_pad_index, :] = x
|
||||
return padded_x
|
||||
|
||||
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
|
||||
results = {
|
||||
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
}
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
|
||||
S = int(variable_block_sizes.sum().item())
|
||||
padded_S = num_blocks * BLOCK_M
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
|
||||
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
|
||||
for _ in range(num_iterations):
|
||||
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
|
||||
# dO_padded = torch.zeros_like(dO_padded)
|
||||
# dO_padded[:, :, non_pad_index, :] = dO
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
|
||||
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
|
||||
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
|
||||
if bs is not None:
|
||||
diff = pt - bs
|
||||
abs_diff = torch.abs(diff)
|
||||
results[name]['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
total_elements = h * S * d * num_iterations
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_graphs(h, d, error_mode='all'):
|
||||
test_configs = [
|
||||
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
|
||||
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
|
||||
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
|
||||
]
|
||||
|
||||
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
|
||||
f"{'gK Avg':<12} {'Rel gK Max':<12} "
|
||||
f"{'gV Avg':<12} {'Rel gV Max':<12} "
|
||||
f"{'gO Avg':<12} {'Rel gO Max':<12}")
|
||||
print("-" * 150)
|
||||
|
||||
for config in test_configs:
|
||||
num_blocks = config["num_blocks"]
|
||||
k = config["k"]
|
||||
description = config["description"]
|
||||
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
|
||||
print(f"{description:<20} {num_blocks:<8} {k:<4} "
|
||||
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
|
||||
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
|
||||
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
|
||||
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
if __name__ == "__main__":
|
||||
h, d = 16, 128
|
||||
print("Block Sparse Attention with Variable Block Sizes Analysis")
|
||||
print("=" * 60)
|
||||
for mode in ['backward']:
|
||||
generate_error_graphs(h, d, error_mode=mode)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch
|
||||
|
||||
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate block sparse mask of shape [h, num_blocks, num_blocks].
|
||||
|
||||
Args:
|
||||
h: number of heads
|
||||
num_blocks: number of blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
"""
|
||||
k = min(k, num_blocks)
|
||||
scores = torch.rand(h, num_blocks, num_blocks, device=device)
|
||||
_, indices = torch.topk(scores, k, dim=-1)
|
||||
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
|
||||
return block_sparse_mask
|
||||
|
||||
|
||||
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
|
||||
"""
|
||||
Convert block-level sparse mask to full attention mask.
|
||||
|
||||
Args:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
variable_block_sizes: [num_blocks] tensor
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
full_mask: [h, S, S] bool tensor where S = total sequence length
|
||||
"""
|
||||
h, num_blocks, _ = block_sparse_mask.shape
|
||||
total_seq_len = variable_block_sizes.sum().item()
|
||||
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
|
||||
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
|
||||
|
||||
for head in range(h):
|
||||
for q_block in range(num_blocks):
|
||||
q_start = cumsum[q_block]
|
||||
q_end = q_start + variable_block_sizes[q_block]
|
||||
|
||||
for kv_block in range(num_blocks):
|
||||
if block_sparse_mask[head, q_block, kv_block]:
|
||||
kv_start = cumsum[kv_block]
|
||||
kv_end = kv_start + variable_block_sizes[kv_block]
|
||||
full_mask[head, q_start:q_end, kv_start:kv_end] = True
|
||||
|
||||
return full_mask
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -0,0 +1,61 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -0,0 +1,15 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'block_sparse': {
|
||||
'source_files': {
|
||||
'h100': 'vsa/block_sparse_h100.cu'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['block_sparse']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -0,0 +1,81 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config_vsa import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.3"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('vsa root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['vsa.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
|
||||
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
Submodule
+1
Submodule csrc/attn/video_sparse_attn/tk added at 6c27e28c81
@@ -0,0 +1,27 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
block_sparse_attn=None
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
|
||||
block_sparse_attn = block_sparse_attn_SM90
|
||||
else:
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_triton
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
block_sparse_attn = block_sparse_attn_triton
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
|
||||
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
NOTE: We assume q, k, v is zero padded!!
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements == 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = (q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
|
||||
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
|
||||
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
|
||||
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
|
||||
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
@@ -0,0 +1,449 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
+78
-91
@@ -672,32 +672,23 @@ block_sparse_attention_forward(
|
||||
torch::Tensor v,
|
||||
torch::Tensor q2k_block_sparse_index,
|
||||
torch::Tensor q2k_block_sparse_num,
|
||||
torch::Tensor kv_block_size
|
||||
torch::Tensor block_size
|
||||
)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
// q shape: (batch, qo_heads, q_seq_len, head_dim)
|
||||
// k shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// v shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// q2k_block_sparse_index shape: (batch, qo_heads, num_q_blocks, max_kv_blocks_per_q)
|
||||
// q2k_block_sparse_num shape: (batch, qo_heads, num_q_blocks)
|
||||
// kv_block_size shape: (num_kv_blocks) This does not need other dimensions because across all batch/heads the padding is the same.
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
|
||||
auto num_q_blocks = q2k_block_sparse_index.size(2);
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
auto num_q_blocks = block_size.size(0);
|
||||
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
|
||||
TORCH_CHECK(num_q_blocks * BLOCK_M == q_seq_len, "This kernel supports variable q block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_kv_blocks * BLOCK_M == kv_seq_len, "This kernel supports variable kv block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -705,8 +696,11 @@ block_sparse_attention_forward(
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(0) == batch, "q2k_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(0) == batch, "q2k_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == num_q_blocks, "q2k_block_sparse_num idx 2 - must match num_q_blocks");
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_index idx 2 - must match seq_len / BLOCK_M");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_num idx 2 - must match seq_len / BLOCK_M");
|
||||
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
@@ -733,12 +727,12 @@ block_sparse_attention_forward(
|
||||
// for the returned outputs
|
||||
torch::Tensor o = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, v.options());
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
@@ -768,11 +762,11 @@ block_sparse_attention_forward(
|
||||
|
||||
using globals = fwd_globals<64>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
|
||||
globals g{
|
||||
qg_arg,
|
||||
@@ -780,17 +774,17 @@ block_sparse_attention_forward(
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<64>,
|
||||
@@ -819,11 +813,11 @@ block_sparse_attention_forward(
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{
|
||||
qg_arg,
|
||||
@@ -831,17 +825,17 @@ block_sparse_attention_forward(
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128>,
|
||||
@@ -868,7 +862,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
torch::Tensor og,
|
||||
torch::Tensor k2q_block_sparse_index,
|
||||
torch::Tensor k2q_block_sparse_num,
|
||||
torch::Tensor kv_block_size)
|
||||
torch::Tensor block_size)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -877,23 +871,11 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
CHECK_INPUT(o);
|
||||
CHECK_INPUT(og);
|
||||
|
||||
// q: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// v: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// o: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// l_vec: [batch, qo_heads, q_seq_len, 1]
|
||||
// og: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k2q_block_sparse_index: [batch, kv_heads, num_kv_blocks, max_num_q_blocks]
|
||||
// k2q_block_sparse_num: [batch, kv_heads, num_kv_blocks]
|
||||
// kv_block_size: [num_kv_blocks]
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index.size(2) must match num_kv_blocks (kv_block_size.size(0))");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -904,18 +886,23 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(0) == batch, "k2q_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(0) == batch, "k2q_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K sequence length");
|
||||
TORCH_CHECK(l_vec.size(2) == q_seq_len, "L sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(o.size(2) == q_seq_len, "O sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(og.size(2) == q_seq_len, "OG sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == num_kv_blocks, "k2q_block_sparse_num idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(l_vec.size(2) == seq_len, "L sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(o.size(2) == seq_len, "O sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(og.size(2) == seq_len, "OG sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_index idx 2 - must match seq_len / BLOCK_N");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_num idx 2 - must match seq_len / BLOCK_N");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(o.size(3) == head_dim, "O head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(og.size(3) == head_dim, "OG head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
|
||||
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
@@ -942,20 +929,20 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
torch::Tensor qg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor kg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor vg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
|
||||
torch::Tensor d_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)}, l_vec.options());
|
||||
|
||||
float* qg_ptr = qg.data_ptr<float>();
|
||||
@@ -984,7 +971,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
// cudaStreamSynchronize(stream);
|
||||
|
||||
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
|
||||
dim3 grid_bwd(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
if (head_dim == 64) {
|
||||
using og_tile = st_bf<4*16, 64>;
|
||||
@@ -997,9 +984,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<64>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1036,15 +1023,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<64>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg,
|
||||
bwd_k_arg,
|
||||
@@ -1055,14 +1042,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
@@ -1101,9 +1088,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<128>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1140,15 +1127,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<128>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg,
|
||||
bwd_k_arg,
|
||||
@@ -1159,14 +1146,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
@@ -1187,4 +1174,4 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
return {qg, kg, vg};
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
|
||||
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
|
||||
from vsa.index import map_to_index
|
||||
from typing import Tuple, Optional
|
||||
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.int()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_triton(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def setup_context_triton(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, M = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
|
||||
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
|
||||
|
||||
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_SM90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
variable_block_sizes = variable_block_sizes.int()
|
||||
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
|
||||
def _block_sparse_attn_SM90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
|
||||
B, H, S, D = q_padded.shape
|
||||
o_padded = torch.empty_like(q_padded)
|
||||
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_SM90(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
|
||||
)
|
||||
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
|
||||
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
|
||||
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
|
||||
return grad_q_padded, grad_k_padded, grad_v_padded
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
|
||||
def _block_sparse_attn_backward_SM90_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch._check(grad_output_padded.dtype == torch.bfloat16)
|
||||
torch._check(lse_padded.dtype == torch.float32)
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_SM90(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
def setup_context_SM90(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, lse_padded = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
|
||||
+1
-4
@@ -1,9 +1,9 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
@@ -26,7 +26,6 @@ def topk_index_to_map_kernel(
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
@@ -60,7 +59,6 @@ def map_to_index_kernel(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
@@ -108,7 +106,6 @@ def topk_index_to_map(index: torch.Tensor,
|
||||
|
||||
return block_map
|
||||
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
@@ -0,0 +1,32 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
|
||||
|
||||
### Installation
|
||||
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
|
||||
|
||||
### Usage
|
||||
|
||||
You can use `moba_attn_varlen` in the following ways:
|
||||
|
||||
**Install from source:**
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
**Import after installation:**
|
||||
```python
|
||||
from vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
**Or import directly from the project root:**
|
||||
```python
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
python csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
```
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from setuptools import find_packages, setup
|
||||
|
||||
PACKAGE_NAME = "vmoba"
|
||||
VERSION = "0.0.0"
|
||||
AUTHOR = "JianzongWu"
|
||||
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
|
||||
URL = "https://github.com/KwaiVGI/VMoBA"
|
||||
|
||||
setup(
|
||||
name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[
|
||||
"flash-attn >= 2.7.1",
|
||||
]
|
||||
)
|
||||
+2
-2
@@ -3,7 +3,7 @@
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from fastvideo_kernel import moba_attn_varlen
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
@@ -51,7 +51,7 @@ def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, dev
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
|
||||
+248
-415
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,195 @@
|
||||
import argparse
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
|
||||
|
||||
def init_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=25)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=8)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
|
||||
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=50)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.1)
|
||||
parser.add_argument("--linear_range", type=float, default=0.75)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def load_model(args):
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
args.shift,
|
||||
args.num_euler_timesteps,
|
||||
linear_quadratic,
|
||||
args.linear_threshold,
|
||||
args.linear_range,
|
||||
)
|
||||
|
||||
if args.transformer_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
# if args.cpu_offload:
|
||||
pipe.enable_sequential_cpu_offload()
|
||||
return pipe
|
||||
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1, )).item()
|
||||
|
||||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||||
|
||||
if not use_negative_prompt:
|
||||
negative_prompt = None
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
output = pipe(
|
||||
prompt=[prompt],
|
||||
negative_prompt=negative_prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=generator,
|
||||
).frames[0]
|
||||
|
||||
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
|
||||
export_to_video(output, output_path, fps=30)
|
||||
return output_path, seed
|
||||
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
args = init_args()
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=args.num_frames,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=args.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=args.num_inference_steps,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
@@ -0,0 +1,15 @@
|
||||
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
|
||||
|
||||
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
|
||||
|
||||
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
|
||||
@@ -1,4 +1,4 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -55,12 +55,18 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -55,12 +55,18 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -55,11 +55,18 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
@@ -55,12 +55,18 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user