Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
404ee8538e | ||
|
|
e57ac59462 | ||
|
|
8c55fdaf7e | ||
|
|
c30779184f | ||
|
|
9d188c0b6c | ||
|
|
9dd7c54221 | ||
|
|
62b95d8287 | ||
|
|
fdf21702f5 | ||
|
|
2972fc9449 |
@@ -7,7 +7,7 @@
|
||||
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.html"><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-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/tMwknPLY" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <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/XcY0Cpv" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -49,10 +49,10 @@ conda activate fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
|
||||
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.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
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:
|
||||
|
||||
@@ -64,7 +64,7 @@ See below for recipes and datasets:
|
||||
|
||||
## 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/video_sparse_attention/installation.html). 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
|
||||
@@ -100,15 +100,15 @@ Run the script with:
|
||||
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.html).
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
|
||||
|
||||
### Other docs:
|
||||
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
- [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.html)
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
|
||||
## 📑 Development Plan
|
||||
@@ -127,7 +127,7 @@ See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/iss
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
|
||||
@@ -247,11 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
|
||||
|
||||
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
|
||||
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install Miniconda and create conda environment
|
||||
COPY Miniconda3-latest-Linux-x86_64.sh /tmp/miniconda.sh
|
||||
RUN bash /tmp/miniconda.sh -b -p /opt/conda && rm /tmp/miniconda.sh
|
||||
ENV PATH=/opt/conda/bin:${PATH}
|
||||
RUN /opt/conda/bin/conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/main
|
||||
RUN /opt/conda/bin/conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/r
|
||||
RUN /opt/conda/bin/conda update -y -n base conda && \
|
||||
/opt/conda/bin/conda create -y -n fastvideo python=3.12 && \
|
||||
/opt/conda/bin/conda clean -afy
|
||||
ENV PATH=/opt/conda/envs/fastvideo/bin:/opt/conda/bin:${PATH}
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Install project dependencies inside the conda environment
|
||||
RUN source /opt/conda/etc/profile.d/conda.sh && \
|
||||
conda activate fastvideo && \
|
||||
pip install --no-cache-dir --upgrade pip && \
|
||||
pip install --no-cache-dir .[dev] && \
|
||||
pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using pip inside conda env and set up shell configuration
|
||||
RUN source /opt/conda/etc/profile.d/conda.sh && \
|
||||
conda activate fastvideo && \
|
||||
pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/conda/etc/profile.d/conda.sh && conda activate fastvideo' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source /opt/conda/etc/profile.d/conda.sh && \
|
||||
conda activate fastvideo && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source /opt/conda/etc/profile.d/conda.sh && \
|
||||
conda activate fastvideo && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
|
||||
ENTRYPOINT ["/bin/bash", "-lc", "source /opt/conda/etc/profile.d/conda.sh && conda activate fastvideo && exec /FastVideo/examples/inference/gradio/start.sh"]
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 122 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 378 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 575 KiB |
@@ -106,6 +106,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
|
||||
return batch
|
||||
```
|
||||
|
||||

|
||||
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
@@ -208,6 +210,10 @@ def step(
|
||||
return prev_sample
|
||||
```
|
||||
|
||||
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
|
||||
|
||||

|
||||
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
@@ -231,6 +237,8 @@ self.attn = LocalAttention(
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
```
|
||||
|
||||

|
||||
|
||||
### Attention Patterns
|
||||
Supports various patterns with memory optimization techniques:
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
num_frame_per_block=7,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -200,7 +200,7 @@ class BaseModelDeployment:
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "fastvideo"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2VModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
|
||||
@@ -210,7 +210,7 @@ class T2VModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fastvideo"}},
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2V14BModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
|
||||
@@ -227,7 +227,7 @@ app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 2})
|
||||
@serve.deployment(num_replicas=50, ray_actor_options={"num_cpus": 2})
|
||||
@serve.ingress(app)
|
||||
class FastVideoAPI:
|
||||
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
python examples/inference/gradio/serving/start_ray_serve_app.py \
|
||||
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers" \
|
||||
--t2v_model_replicas "1"
|
||||
python examples/inference/gradio/start_ray_serve_app.py \
|
||||
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
|
||||
--t2v_model_replicas "4,4"
|
||||
@@ -14,7 +14,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig)
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -38,6 +38,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
|
||||
@@ -176,3 +176,13 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
|
||||
is_causal: bool = True
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
@@ -11,7 +11,7 @@ from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_Sampl
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480PConfig,
|
||||
FastWanT2V480P_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
@@ -21,7 +21,8 @@ from fastvideo.configs.sample.wan import (
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -64,7 +65,7 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480PConfig,
|
||||
FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
@@ -72,11 +73,16 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
SelfForcingWanT2V480PConfig,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.2
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
|
||||
# Cosmos2
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World":
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -86,6 +92,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -96,7 +104,9 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam
|
||||
"wandmdpipeline": FastWanT2V480P_SamplingParam,
|
||||
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# DMD parameters
|
||||
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
|
||||
num_inference_steps: int = 3
|
||||
@@ -183,5 +183,18 @@ class Wan2_2_Fun_A14B_Control_SamplingParam(
|
||||
# ============= Causal Self-Forcing =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
|
||||
Wan2_2_T2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 2.0
|
||||
guidance_scale_2: float = 2.0
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
# Copyright 2024 The TorchTune Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
from __future__ import annotations
|
||||
import os
|
||||
import contextlib
|
||||
from collections.abc import Callable, Generator
|
||||
from itertools import chain
|
||||
@@ -197,10 +199,11 @@ def shard_model(
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
|
||||
logger.warning(
|
||||
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
|
||||
type(model).__name__)
|
||||
# Check if we should use size-based filtering
|
||||
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
|
||||
|
||||
if not fsdp_shard_conditions:
|
||||
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
|
||||
return
|
||||
|
||||
fsdp_kwargs = {
|
||||
@@ -215,20 +218,38 @@ def shard_model(
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
# TODO(will): don't reshard after forward for the last layer to save on the
|
||||
# all-gather that will immediately happen Shard the model with FSDP,
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
for shard_condition in fsdp_shard_conditions
|
||||
]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
if use_size_filtering:
|
||||
# Size-based filtering mode
|
||||
min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000"))
|
||||
logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6)
|
||||
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
# Count all parameters
|
||||
param_count = sum(p.numel() for p in m.parameters(recurse=True))
|
||||
|
||||
# Skip small modules
|
||||
if param_count < min_params:
|
||||
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)",
|
||||
n, param_count / 1e6, min_params / 1e6)
|
||||
continue
|
||||
|
||||
# Shard this module
|
||||
logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6)
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
else:
|
||||
# Shard all modules matching conditions
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||
@@ -729,9 +729,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(timestep, src=0)
|
||||
|
||||
timestep = shift_timestep(
|
||||
timestep,
|
||||
@@ -844,9 +841,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(fake_score_timestep, src=0)
|
||||
|
||||
fake_score_timestep = shift_timestep(
|
||||
fake_score_timestep,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import copy
|
||||
import os
|
||||
import time
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
@@ -45,6 +46,13 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the self-forcing training pipeline."""
|
||||
# Check if FSDP2 auto wrap is enabled - not supported for self-forcing distillation
|
||||
if os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1":
|
||||
raise NotImplementedError(
|
||||
"FASTVIDEO_FSDP2_AUTOWRAP is not implemented for self-forcing distillation. "
|
||||
"Please set FASTVIDEO_FSDP2_AUTOWRAP=0 or unset the environment variable."
|
||||
)
|
||||
|
||||
logger.info("Initializing self-forcing distillation pipeline...")
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
|
||||
@@ -467,6 +467,10 @@ class WorkerMultiprocProc:
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(
|
||||
method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
else:
|
||||
result = self.worker.execute_method(method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
|
||||
@@ -14,6 +14,7 @@ edit_uri: edit/main/docs/
|
||||
# Configuration
|
||||
theme:
|
||||
name: material
|
||||
favicon: assets/logos/icon_simple.svg
|
||||
palette:
|
||||
- scheme: default
|
||||
toggle:
|
||||
|
||||
Reference in New Issue
Block a user