Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
350afecede | ||
|
|
c868206c40 | ||
|
|
5a4eebdf97 | ||
|
|
f859409844 | ||
|
|
ed5c2e0cc0 | ||
|
|
25a91bfb0f |
+21
-21
@@ -16,9 +16,9 @@ steps:
|
||||
diff: 'git fetch origin "$BUILDKITE_PULL_REQUEST_BASE_BRANCH" && git diff --name-only origin/"$BUILDKITE_PULL_REQUEST_BASE_BRANCH"...HEAD'
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
- "fastvideo/models/encoders/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/encoders/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -29,9 +29,9 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/vaes/**"
|
||||
- "fastvideo/models/vaes/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/vaes/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -42,11 +42,11 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/attention/**"
|
||||
- "fastvideo/models/dits/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/attention/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -57,7 +57,7 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**/*.py"
|
||||
- "fastvideo/**/*.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -68,11 +68,11 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/tests/lora/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/pipelines/**"
|
||||
- "fastvideo/v1/layers/lora/**"
|
||||
- "fastvideo/tests/lora/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/pipelines/**"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -83,7 +83,7 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -94,7 +94,7 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -105,7 +105,7 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
@@ -121,7 +121,7 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
|
||||
@@ -50,7 +50,7 @@ else
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
|
||||
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
|
||||
@@ -23,7 +23,7 @@ body:
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
Please share your environment with us. You can run the command **python collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -8,14 +8,14 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
@@ -115,37 +115,37 @@ jobs:
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
@@ -167,7 +167,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -185,7 +185,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -203,7 +203,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -229,7 +229,7 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -248,7 +248,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -268,7 +268,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -288,7 +288,7 @@ jobs:
|
||||
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/v1/tests/inference/STA -srP"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -343,7 +343,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
|
||||
+1
-1
@@ -57,7 +57,7 @@ docs/source/training/examples/
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
|
||||
@@ -3,7 +3,7 @@ default_stages:
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/v1/third_party/.*|
|
||||
fastvideo/third_party/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
@@ -69,7 +69,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep -v "^fastvideo/v1/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
|
||||
|
||||
@@ -16,7 +16,7 @@ import sys
|
||||
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
|
||||
from collections import namedtuple
|
||||
|
||||
from fastvideo.v1.envs import environment_variables
|
||||
from fastvideo.envs import environment_variables
|
||||
|
||||
try:
|
||||
import torch
|
||||
@@ -81,6 +81,7 @@ DEFAULT_CONDA_PATTERNS = {
|
||||
DEFAULT_PIP_PATTERNS = {
|
||||
"torch",
|
||||
"numpy",
|
||||
"mypy",
|
||||
"flake8",
|
||||
"triton",
|
||||
"optree",
|
||||
+2
-2
@@ -74,8 +74,8 @@ After installation, the following nodes will be available in the ComfyUI interfa
|
||||
You may have noticed many arguments on the nodes have 'auto' as the default value. This is because FastVideo will automatically detect the best values for these parameters based on the model and the hardware. However, you can also manually configure these parameters to get the best performance for your specific use case. We plan on releasing more optimized workflow files for different models and hardware configurations in the future.
|
||||
|
||||
You can see what some of the default configurations are by looking at the FastVideo repo:
|
||||
- [Wan2.1-I2V-14B-480P-Diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/v1/configs/wan_14B_i2v_480p_pipeline.json)
|
||||
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/v1/configs/fasthunyuan_t2v.json)
|
||||
- [Wan2.1-I2V-14B-480P-Diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/wan_14B_i2v_480p_pipeline.json)
|
||||
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/fasthunyuan_t2v.json)
|
||||
|
||||
### Node Configuration
|
||||
|
||||
|
||||
@@ -9,11 +9,11 @@
|
||||
## Initialization Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.pipelines.PipelineConfig
|
||||
fastvideo.configs.pipelines.PipelineConfig
|
||||
```
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.sample.SamplingParam
|
||||
fastvideo.configs.sample.SamplingParam
|
||||
```
|
||||
|
||||
@@ -1,24 +1,24 @@
|
||||
# 🔍 FastVideo Overview
|
||||
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/v1/` codebase.
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/` codebase.
|
||||
|
||||
## Table of Contents - V1 Directory Structure and Files
|
||||
## Table of Contents - Directory Structure and Files
|
||||
|
||||
- [`fastvideo/v1/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/v1/models/`](#design-model-components) - Model implementations
|
||||
- [`fastvideo/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/models/`](#design-model-components) - Model implementations
|
||||
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/v1/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/v1/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/v1/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/v1/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/v1/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/v1/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/v1/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/v1/utils.py` - Utility functions
|
||||
- [`fastvideo/v1/logger.py`](#design-logger) - Logging infrastructure
|
||||
- [`fastvideo/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/utils.py` - Utility functions
|
||||
- [`fastvideo/logger.py`](#design-logger) - Logging infrastructure
|
||||
|
||||
## Core Architecture
|
||||
|
||||
@@ -32,7 +32,7 @@ FastVideo separates model components from execution logic with these principles:
|
||||
(design-fastvideo-args)=
|
||||
## FastVideoArgs
|
||||
|
||||
The `FastVideoArgs` class in `fastvideo/v1/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
|
||||
Key features include:
|
||||
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
|
||||
@@ -111,7 +111,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
|
||||
(design-forwardbatch)=
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
|
||||
- **Input Data**: Prompts, images, generation parameters
|
||||
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
|
||||
@@ -123,14 +123,14 @@ This structure facilitates clear state transitions between stages.
|
||||
(design-model-components)=
|
||||
## Model Components
|
||||
|
||||
The `fastvideo/v1/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
|
||||
(design-transformer-models)=
|
||||
### Transformer Models
|
||||
|
||||
Transformer networks perform the actual denoising during diffusion:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/dits/`
|
||||
- **Location**: `fastvideo/models/dits/`
|
||||
- **Examples**:
|
||||
- `WanTransformer3DModel`
|
||||
- `HunyuanVideoTransformer3DModel`
|
||||
@@ -157,7 +157,7 @@ def forward(
|
||||
|
||||
VAEs handle conversion between pixel space and latent space:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/vaes/`
|
||||
- **Location**: `fastvideo/models/vaes/`
|
||||
- **Examples**:
|
||||
- `AutoencoderKLWan`
|
||||
- `AutoencoderKLHunyuanVideo`
|
||||
@@ -175,7 +175,7 @@ FastVideo's VAE implementations include:
|
||||
|
||||
Encoders process conditioning inputs into embeddings:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/encoders/`
|
||||
- **Location**: `fastvideo/models/encoders/`
|
||||
- **Text Encoders**:
|
||||
- `CLIPTextModel`
|
||||
- `LlamaModel`
|
||||
@@ -193,7 +193,7 @@ FastVideo implements optimizations such as:
|
||||
|
||||
Schedulers manage the diffusion sampling process:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/schedulers/`
|
||||
- **Location**: `fastvideo/models/schedulers/`
|
||||
- **Examples**:
|
||||
- `UniPCMultistepScheduler`
|
||||
- `FlowMatchEulerDiscreteScheduler`
|
||||
@@ -219,7 +219,7 @@ def step(
|
||||
(design-optimized-attention)=
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/v1/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
|
||||
### Attention Backends
|
||||
Multiple implementations with automatic selection:
|
||||
@@ -248,7 +248,7 @@ Supports various patterns with memory optimization techniques:
|
||||
(design-distributed-processing)=
|
||||
## Distributed Processing
|
||||
|
||||
The `fastvideo/v1/distributed/` directory contains implementations for distributed model execution:
|
||||
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
|
||||
|
||||
(design-tensor-parallelism)=
|
||||
### Tensor Parallelism
|
||||
@@ -260,7 +260,7 @@ Tensor parallelism splits model weights across devices:
|
||||
|
||||
```python
|
||||
# Tensor-parallel layers in a transformer block
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
|
||||
# Split along output dimension
|
||||
self.qkv_proj = ColumnParallelLinear(
|
||||
@@ -288,7 +288,7 @@ Sequence parallelism splits sequences across devices:
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
@@ -312,7 +312,7 @@ Efficient communication primitives minimize distributed overhead:
|
||||
|
||||
### ForwardContext
|
||||
|
||||
Defined in `fastvideo/v1/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
|
||||
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
@@ -333,7 +333,7 @@ with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
|
||||
(design-executor-and-worker-abstractions)=
|
||||
## Executor and Worker System
|
||||
|
||||
The `fastvideo/v1/worker/` directory contains the distributed execution framework:
|
||||
The `fastvideo/worker/` directory contains the distributed execution framework:
|
||||
|
||||
### Executor Abstraction
|
||||
|
||||
@@ -360,7 +360,7 @@ This design allows FastVideo to efficiently utilize multiple GPUs while providin
|
||||
(design-platforms)=
|
||||
## Platforms
|
||||
|
||||
The `fastvideo/v1/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
|
||||
### Platform Abstraction
|
||||
|
||||
@@ -377,7 +377,7 @@ The primary components include:
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.platforms import current_platform, _Backend
|
||||
from fastvideo.platforms import current_platform, _Backend
|
||||
|
||||
# Check hardware capabilities
|
||||
if current_platform.supports_backend(_Backend.FLASH_ATTN):
|
||||
@@ -398,9 +398,9 @@ See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
If you're a new contributor, here are some common areas to explore:
|
||||
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/v1/models/`
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/models/`
|
||||
2. **Optimizing performance**: Look at attention implementations or memory management
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/v1/pipelines/`
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/pipelines/`
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
|
||||
@@ -10,20 +10,20 @@ This class will be the primary Python API for generating videos and images.
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.v1.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.v1.worker.executor.Executor], log_stats: bool)
|
||||
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.worker.executor.Executor], log_stats: bool)
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
|
||||
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.v1.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -38,25 +38,25 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
|
||||
```
|
||||
|
||||
`````{py:class} PipelineConfig
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
|
||||
````{py:method} dump_to_json(file_path: str)
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -68,16 +68,16 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
|
||||
```
|
||||
|
||||
`````{py:class} SamplingParam
|
||||
:canonical: fastvideo.v1.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -46,25 +46,25 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
- Encoders: `fastvideo/v1/models/encoders/`
|
||||
- VAEs: `fastvideo/v1/models/vaes/`
|
||||
- Transformer models: `fastvideo/v1/models/dits/`
|
||||
- Schedulers: `fastvideo/v1/models/schedulers/`
|
||||
- Encoders: `fastvideo/models/encoders/`
|
||||
- VAEs: `fastvideo/models/vaes/`
|
||||
- Transformer models: `fastvideo/models/dits/`
|
||||
- Schedulers: `fastvideo/models/schedulers/`
|
||||
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
- nn.LayerNorm → fastvideo.v1.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.v1.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.v1.layers.activation
|
||||
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
# Output dimension parallelism
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear
|
||||
from fastvideo.layers.linear import ColumnParallelLinear
|
||||
self.q_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=head_size * num_heads,
|
||||
@@ -73,7 +73,7 @@ self.q_proj = ColumnParallelLinear(
|
||||
)
|
||||
|
||||
# Fused QKV projection
|
||||
from fastvideo.v1.layers.linear import QKVParallelLinear
|
||||
from fastvideo.layers.linear import QKVParallelLinear
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=attention_head_dim,
|
||||
@@ -82,7 +82,7 @@ self.qkv_proj = QKVParallelLinear(
|
||||
)
|
||||
|
||||
# Input dimension parallelism
|
||||
from fastvideo.v1.layers.linear import RowParallelLinear
|
||||
from fastvideo.layers.linear import RowParallelLinear
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=head_size * num_heads,
|
||||
output_size=hidden_size,
|
||||
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -130,7 +130,7 @@ self.attn = DistributedAttention(
|
||||
Register implemented modules in the model registry:
|
||||
|
||||
```python
|
||||
# In fastvideo/v1/models/registry.py
|
||||
# In fastvideo/models/registry.py
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
|
||||
}
|
||||
@@ -145,7 +145,7 @@ _VAE_MODELS = {
|
||||
Create a new directory for your pipeline:
|
||||
|
||||
```
|
||||
fastvideo/v1/pipelines/
|
||||
fastvideo/pipelines/
|
||||
├── your_pipeline/
|
||||
│ ├── __init__.py
|
||||
│ └── your_pipeline.py
|
||||
@@ -167,13 +167,13 @@ Pipelines are composed of stages, each handling a specific part of the diffusion
|
||||
### Creating Your Pipeline
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (
|
||||
InputValidationStage, CLIPTextEncodingStage, TimestepPreparationStage,
|
||||
LatentPreparationStage, DenoisingStage, DecodingStage
|
||||
)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import torch
|
||||
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
@@ -246,7 +246,7 @@ EntryClass = MyCustomPipeline
|
||||
If existing stages don't meet your needs, create custom ones:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
class MyCustomStage(PipelineStage):
|
||||
"""Custom processing stage for the pipeline."""
|
||||
@@ -305,7 +305,7 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
1. Scan all packages under `fastvideo/v1/pipelines/`
|
||||
1. Scan all packages under `fastvideo/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
def main():
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
def main():
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig, SamplingParam
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_fp16"
|
||||
def main():
|
||||
|
||||
@@ -6,7 +6,7 @@ import gradio as gr
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Inference using a LoRA checkpoint from FastVideo trainer.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -84,7 +84,7 @@ miscellaneous_args=(
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -120,7 +120,7 @@ srun torchrun \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
+1
-1
@@ -7,7 +7,7 @@ DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v_1_3b_inp/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
|
||||
@@ -84,7 +84,7 @@ miscellaneous_args=(
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -120,7 +120,7 @@ srun torchrun \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -83,7 +83,7 @@ miscellaneous_args=(
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -7,7 +7,7 @@ DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
|
||||
@@ -83,7 +83,7 @@ miscellaneous_args=(
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -117,7 +117,7 @@ srun torchrun \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -83,7 +83,7 @@ torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port 29501 \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
|
||||
@@ -7,7 +7,7 @@ DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
@@ -0,0 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.attention.layer import (DistributedAttention,
|
||||
DistributedAttention_VSA, LocalAttention)
|
||||
from fastvideo.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
+2
-2
@@ -6,8 +6,8 @@ from dataclasses import dataclass, fields
|
||||
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
import torch
|
||||
|
||||
+5
-5
@@ -12,11 +12,11 @@ try:
|
||||
except ImportError:
|
||||
flash_attn_func = flash_attn_2_func
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+3
-5
@@ -3,11 +3,9 @@
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -2,11 +2,9 @@
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+11
-11
@@ -7,17 +7,17 @@ import torch
|
||||
from einops import rearrange
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+8
-8
@@ -10,14 +10,14 @@ try:
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -3,15 +3,14 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention.selector import (backend_name_to_enum,
|
||||
get_attn_backend)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
from fastvideo.attention.selector import backend_name_to_enum, get_attn_backend
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
@@ -9,11 +9,11 @@ from typing import cast
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
@@ -2,7 +2,7 @@
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig"]
|
||||
@@ -2,9 +2,9 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
+1
-1
@@ -3,7 +3,7 @@ from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_double_block(n: str, m) -> bool:
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
@@ -0,0 +1,14 @@
|
||||
from fastvideo.configs.models.encoders.base import (BaseEncoderOutput,
|
||||
EncoderConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig,
|
||||
CLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
|
||||
"T5Config"
|
||||
]
|
||||
+3
-3
@@ -4,9 +4,9 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
+4
-4
@@ -1,10 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
+2
-2
@@ -1,8 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
+2
-2
@@ -1,8 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
@@ -0,0 +1,9 @@
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVAEConfig",
|
||||
"WanVAEConfig",
|
||||
"StepVideoVAEConfig",
|
||||
]
|
||||
@@ -6,8 +6,8 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.utils import StoreBoolean
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.utils import StoreBoolean
|
||||
|
||||
|
||||
@dataclass
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
+1
-1
@@ -3,7 +3,7 @@ from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -0,0 +1,15 @@
|
||||
from fastvideo.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
@@ -7,13 +7,12 @@ from typing import Any, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.v1.configs.utils import update_config_from_args
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
shallow_asdict)
|
||||
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.utils import update_config_from_args
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean, shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -213,11 +212,11 @@ class PipelineConfig:
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||
|
||||
return parser
|
||||
@@ -241,7 +240,7 @@ class PipelineConfig:
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
@@ -256,7 +255,7 @@ class PipelineConfig:
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
@@ -5,12 +5,12 @@ from typing import TypedDict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig, LlamaConfig)
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig, LlamaConfig)
|
||||
from fastvideo.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
+9
-12
@@ -4,18 +4,15 @@
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import (FastWanT2V480PConfig,
|
||||
WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (FastWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+4
-4
@@ -1,10 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import StepVideoConfig
|
||||
from fastvideo.v1.configs.models.vaes import StepVideoVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import StepVideoConfig
|
||||
from fastvideo.configs.models.vaes import StepVideoVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -4,12 +4,12 @@ from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPVisionConfig, T5Config)
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPVisionConfig, T5Config)
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
__all__ = ["SamplingParam"]
|
||||
@@ -2,7 +2,7 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -66,7 +66,7 @@ class SamplingParam:
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.v1.configs.sample.registry import (
|
||||
from fastvideo.configs.sample.registry import (
|
||||
get_sampling_param_cls_for_name)
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
if sampling_cls is not None:
|
||||
@@ -1,8 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.teacache import TeaCacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -3,17 +3,17 @@ import os
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.v1.configs.sample.wan import (FastWanT2V480PConfig,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.configs.sample.wan import (FastWanT2V480PConfig,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import CacheParams
|
||||
from fastvideo.configs.sample.base import CacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -1,8 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.configs.sample.teacache import WanTeaCacheParams
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.teacache import WanTeaCacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -2,13 +2,12 @@
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import (
|
||||
VideoCaptionMergedDataset)
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
|
||||
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
+7
-7
@@ -7,13 +7,13 @@ import time
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
|
||||
from fastvideo.dataset.parquet_dataset_iterable_style import (
|
||||
build_parquet_iterable_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
from fastvideo.distributed import get_world_rank
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -52,9 +52,9 @@ def main() -> None:
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
+8
-8
@@ -8,14 +8,14 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
from fastvideo.distributed import get_world_rank
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -55,9 +55,9 @@ def main() -> None:
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
+4
-4
@@ -9,10 +9,10 @@ import tqdm
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.dataset.utils import collate_latents_embs_masks
|
||||
from fastvideo.distributed import (get_sp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+4
-4
@@ -13,10 +13,10 @@ import tqdm
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
from fastvideo.distributed import (get_sp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+1
-1
@@ -16,7 +16,7 @@ from einops import rearrange
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+4
-4
@@ -6,10 +6,10 @@ import pathlib
|
||||
import datasets
|
||||
from torch.utils.data import IterableDataset
|
||||
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image, load_video
|
||||
from fastvideo.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import load_image, load_video
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
from fastvideo.distributed.communication_op import *
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
|
||||
get_local_torch_device, get_sp_group, get_sp_parallel_rank,
|
||||
get_sp_world_size, get_tp_group, get_tp_rank, get_tp_world_size,
|
||||
@@ -9,7 +9,7 @@ from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
from fastvideo.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
# Initialization
|
||||
+1
-1
@@ -4,7 +4,7 @@
|
||||
import torch
|
||||
import torch.distributed
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_group, get_tp_group
|
||||
from fastvideo.distributed.parallel_state import get_sp_group, get_tp_group
|
||||
|
||||
|
||||
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
+2
-2
@@ -6,8 +6,8 @@ import os
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.v1.platforms.interface import CpuArchEnum
|
||||
from fastvideo.platforms import current_platform
|
||||
from fastvideo.platforms.interface import CpuArchEnum
|
||||
|
||||
from .base_device_communicator import DeviceCommunicatorBase
|
||||
|
||||
+2
-2
@@ -4,7 +4,7 @@
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
from fastvideo.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase)
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
unique_name: str = ""):
|
||||
super().__init__(cpu_group, device, device_group, unique_name)
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl import (
|
||||
from fastvideo.distributed.device_communicators.pynccl import (
|
||||
PyNcclCommunicator)
|
||||
|
||||
self.pynccl_comm: PyNcclCommunicator | None = None
|
||||
+4
-4
@@ -6,12 +6,12 @@ import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl_wrapper import (
|
||||
from fastvideo.distributed.device_communicators.pynccl_wrapper import (
|
||||
NCCLLibrary, buffer_type, cudaStream_t, ncclComm_t, ncclDataTypeEnum,
|
||||
ncclRedOpTypeEnum, ncclUniqueId)
|
||||
from fastvideo.v1.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import current_stream
|
||||
from fastvideo.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import current_stream
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+2
-2
@@ -32,8 +32,8 @@ from typing import Any
|
||||
import torch
|
||||
from torch.distributed import ReduceOp
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import find_nccl_library
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import find_nccl_library
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
+11
-11
@@ -38,14 +38,14 @@ import torch
|
||||
import torch.distributed
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase)
|
||||
from fastvideo.v1.distributed.device_communicators.cpu_communicator import (
|
||||
from fastvideo.distributed.device_communicators.cpu_communicator import (
|
||||
CpuCommunicator)
|
||||
from fastvideo.v1.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -184,7 +184,7 @@ class GroupCoordinator:
|
||||
print(f"rank: {self.rank} group not found")
|
||||
raise e
|
||||
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
# TODO: fix it for other platforms
|
||||
self.device = get_local_torch_device()
|
||||
@@ -195,7 +195,7 @@ class GroupCoordinator:
|
||||
if use_device_communicator and self.world_size > 1:
|
||||
# Platform-aware device communicator selection
|
||||
if current_platform.is_cuda_alike():
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
from fastvideo.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
self.device_communicator = CudaCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
@@ -214,7 +214,7 @@ class GroupCoordinator:
|
||||
|
||||
self.mq_broadcaster = None
|
||||
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
# TODO(will): check if this is needed
|
||||
# self.use_custom_op_call = current_platform.is_cuda_alike()
|
||||
@@ -258,7 +258,7 @@ class GroupCoordinator:
|
||||
def graph_capture(self,
|
||||
graph_capture_context: GraphCaptureContext | None = None):
|
||||
# Platform-aware graph capture
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
if graph_capture_context is None:
|
||||
@@ -775,7 +775,7 @@ def init_distributed_environment(
|
||||
device_id: torch.device | None = None,
|
||||
):
|
||||
# Determine the appropriate backend based on the platform
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.platforms import current_platform
|
||||
if backend == "nccl" and not current_platform.is_cuda_alike():
|
||||
# Use gloo backend for non-CUDA platforms (MPS, CPU)
|
||||
backend = "gloo"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user