Compare commits

..
19 changed files with 161 additions and 128 deletions
+9 -9
View File
@@ -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)
-80
View File
@@ -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

+8
View File
@@ -106,6 +106,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
return batch
```
![Pipeline execution and data flow](../assets/images/pipeline.png)
### 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.
![Model loading flow](../assets/images/load_models.png)
## 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 backend selector design](../assets/images/attention_backend.png)
### 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:
+3 -3
View File
@@ -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"
+2 -1
View File
@@ -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,
+10
View File
@@ -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
+15 -5
View File
@@ -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
}
+15 -2
View File
@@ -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
+39 -18
View File
@@ -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
+4
View File
@@ -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)
+1
View File
@@ -14,6 +14,7 @@ edit_uri: edit/main/docs/
# Configuration
theme:
name: material
favicon: assets/logos/icon_simple.svg
palette:
- scheme: default
toggle: