Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
92fbe6e745 | ||
|
|
ec53a4dce7 | ||
|
|
d0ccf2d83e | ||
|
|
dfa226f37e | ||
|
|
6c10900e6e | ||
|
|
54d825e7de | ||
|
|
0fb82f071f | ||
|
|
278e52d0dc | ||
|
|
be62cfaa52 | ||
|
|
f845be4a0e | ||
|
|
e8b85b37e9 | ||
|
|
33bd6b4d33 |
@@ -0,0 +1 @@
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
|
||||
Type, TypeVar)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
import torch
|
||||
@@ -154,7 +154,7 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: "ForwardBatch",
|
||||
inference_args: "InferenceArgs",
|
||||
fastvideo_args: "FastVideoArgs",
|
||||
) -> T:
|
||||
"""Build attention metadata with on-device tensors."""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -12,7 +12,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -77,7 +77,7 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig"
|
||||
]
|
||||
@@ -0,0 +1,70 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Dict, Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 125
|
||||
fps: int = 24
|
||||
|
||||
# Video generation parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
seed: int = 1024
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1
|
||||
hidden_state_skip_layer: int = 0
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
# Additional parameters can be added as a dict
|
||||
extra_params: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
# Add sliding tile specific parameters
|
||||
window_size: int = 16
|
||||
stride: int = 8
|
||||
|
||||
# You can provide custom defaults for inherited fields
|
||||
height: int = 576
|
||||
width: int = 1024
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
@@ -0,0 +1,39 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanConfig(BaseConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp16"
|
||||
|
||||
# HunyuanConfig-specific added parameters
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanConfig(HunyuanConfig):
|
||||
"""Configuration specifically optimized for FastHunyuan weights."""
|
||||
|
||||
# Override HunyuanConfig defaults
|
||||
num_inference_steps: int = 6
|
||||
flow_shift: int = 17
|
||||
|
||||
# No need to re-specify guidance_scale or embedded_cfg_scale as they
|
||||
# already have the desired values from HunyuanConfig
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
from typing import Dict, Type, Optional, Callable
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
|
||||
from fastvideo.v1.utils import maybe_download_model_index, verify_model_config_and_directory
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
|
||||
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[Type[BaseConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
|
||||
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
print(pipeline_name)
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
|
||||
break
|
||||
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
@@ -0,0 +1,44 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(BaseConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
neg_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
flow_shift: int = 3
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
text_len: int = 512
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
# Precision for each component
|
||||
image_encoder_precision: str = "fp32"
|
||||
@@ -6,7 +6,7 @@ from typing import List, cast
|
||||
|
||||
from fastvideo.v1.entrypoints.cli import utils
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
@@ -82,7 +82,7 @@ class GenerateSubcommand(CLISubcommand):
|
||||
default=None,
|
||||
help="Port for the master process")
|
||||
|
||||
generate_parser = InferenceArgs.add_cli_args(generate_parser)
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
|
||||
@@ -4,23 +4,33 @@
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class InferenceArgs:
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
# Distributed executor backend
|
||||
distributed_executor_backend: str = "torch"
|
||||
|
||||
inference_mode: bool = True # if False == training mode
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: Optional[str] = None
|
||||
|
||||
# Parallelism
|
||||
tp_size: int = 1
|
||||
sp_size: int = 1
|
||||
num_gpus: int = 1
|
||||
tp_size: Optional[int] = None
|
||||
sp_size: Optional[int] = None
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
@@ -112,40 +122,55 @@ class InferenceArgs:
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
|
||||
# distributed_executor_backend
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
choices=["mp", "ray", "torch"],
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
action="store_true",
|
||||
default=InferenceArgs.trust_remote_code,
|
||||
default=FastVideoArgs.trust_remote_code,
|
||||
help="Trust remote code when loading HuggingFace models",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--revision",
|
||||
type=str,
|
||||
default=InferenceArgs.revision,
|
||||
default=FastVideoArgs.revision,
|
||||
help=
|
||||
"The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
)
|
||||
|
||||
# Parallelism
|
||||
parser.add_argument(
|
||||
"--num-gpus",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_gpus,
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.tp_size,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.sp_size,
|
||||
default=FastVideoArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dist-timeout",
|
||||
type=int,
|
||||
default=InferenceArgs.dist_timeout,
|
||||
default=FastVideoArgs.dist_timeout,
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
@@ -153,56 +178,56 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=InferenceArgs.height,
|
||||
default=FastVideoArgs.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=InferenceArgs.width,
|
||||
default=FastVideoArgs.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=InferenceArgs.num_frames,
|
||||
default=FastVideoArgs.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=InferenceArgs.num_inference_steps,
|
||||
default=FastVideoArgs.num_inference_steps,
|
||||
help="Number of inference steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_scale,
|
||||
default=FastVideoArgs.guidance_scale,
|
||||
help="Guidance scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_rescale,
|
||||
default=FastVideoArgs.guidance_rescale,
|
||||
help="Guidance rescale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.embedded_cfg_scale,
|
||||
default=FastVideoArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=float,
|
||||
default=InferenceArgs.flow_shift,
|
||||
default=FastVideoArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
default=InferenceArgs.output_type,
|
||||
default=FastVideoArgs.output_type,
|
||||
choices=["pil"],
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
@@ -210,7 +235,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=InferenceArgs.precision,
|
||||
default=FastVideoArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
@@ -219,14 +244,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.vae_precision,
|
||||
default=FastVideoArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action="store_true",
|
||||
default=InferenceArgs.vae_tiling,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -238,14 +263,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision,
|
||||
default=FastVideoArgs.text_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len,
|
||||
default=FastVideoArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
|
||||
@@ -253,7 +278,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.image_encoder_precision,
|
||||
default=FastVideoArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
@@ -263,14 +288,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision_2,
|
||||
default=FastVideoArgs.text_encoder_precision_2,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for secondary text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len_2,
|
||||
default=FastVideoArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
@@ -278,13 +303,13 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default=InferenceArgs.flow_solver,
|
||||
default=FastVideoArgs.flow_solver,
|
||||
help="Solver for flow matching",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default=InferenceArgs.denoise_type,
|
||||
default=FastVideoArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
@@ -305,7 +330,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
type=str,
|
||||
default=InferenceArgs.scheduler_type,
|
||||
default=FastVideoArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
@@ -313,19 +338,19 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
type=str,
|
||||
default=InferenceArgs.neg_prompt,
|
||||
default=FastVideoArgs.neg_prompt,
|
||||
help="Negative prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos",
|
||||
type=int,
|
||||
default=InferenceArgs.num_videos,
|
||||
default=FastVideoArgs.num_videos,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=InferenceArgs.fps,
|
||||
default=FastVideoArgs.fps,
|
||||
help="Frames per second for output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -344,7 +369,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=InferenceArgs.log_level,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
@@ -368,20 +393,20 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=InferenceArgs.output_path,
|
||||
default=FastVideoArgs.output_path,
|
||||
help="Directory to save generated videos",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=InferenceArgs.seed,
|
||||
default=FastVideoArgs.seed,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
@@ -408,6 +433,15 @@ class InferenceArgs:
|
||||
|
||||
def check_inference_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if self.tp_size is None:
|
||||
self.tp_size = self.num_gpus
|
||||
if self.sp_size is None:
|
||||
self.sp_size = self.num_gpus
|
||||
|
||||
if self.tp_size != self.sp_size:
|
||||
raise ValueError(
|
||||
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
|
||||
)
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
@@ -418,10 +452,10 @@ class InferenceArgs:
|
||||
raise ValueError("prompt_path must be a text file")
|
||||
|
||||
|
||||
_inference_args = None
|
||||
_current_fastvideo_args = None
|
||||
|
||||
|
||||
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
|
||||
@@ -433,26 +467,38 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
The inference arguments.
|
||||
"""
|
||||
parser = FlexibleArgumentParser()
|
||||
InferenceArgs.add_cli_args(parser)
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
inference_args = InferenceArgs.from_cli_args(raw_args)
|
||||
inference_args.check_inference_args()
|
||||
global _inference_args
|
||||
_inference_args = inference_args
|
||||
return inference_args
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
|
||||
fastvideo_args.check_inference_args()
|
||||
global _current_fastvideo_args
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
return fastvideo_args
|
||||
|
||||
|
||||
def get_inference_args() -> InferenceArgs:
|
||||
global _inference_args
|
||||
if _inference_args is None:
|
||||
raise ValueError("Inference arguments not set")
|
||||
return _inference_args
|
||||
@contextmanager
|
||||
def set_current_fastvideo_args(fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Temporarily set the current fastvideo config.
|
||||
Used during model initialization.
|
||||
We save the current fastvideo config in a global variable,
|
||||
so that all modules can access it, e.g. custom ops
|
||||
can access the fastvideo config to determine how to dispatch.
|
||||
"""
|
||||
global _current_fastvideo_args
|
||||
old_fastvideo_args = _current_fastvideo_args
|
||||
try:
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
yield
|
||||
finally:
|
||||
_current_fastvideo_args = old_fastvideo_args
|
||||
|
||||
|
||||
class DeprecatedAction(argparse.Action):
|
||||
|
||||
def __init__(self, option_strings, dest, nargs=0, **kwargs):
|
||||
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise ValueError(self.help)
|
||||
def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
if _current_fastvideo_args is None:
|
||||
# in ci, usually when we test custom ops/modules directly,
|
||||
# we don't set the fastvideo config. In that case, we set a default
|
||||
# config.
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
@@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -52,7 +52,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
inference_args: Optional[InferenceArgs] = None):
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
Here we can inject common logic for every model forward pass.
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
build_pipeline)
|
||||
@@ -28,29 +28,29 @@ class InferenceEngine:
|
||||
def __init__(
|
||||
self,
|
||||
pipeline: ComposedPipelineBase,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
):
|
||||
"""
|
||||
Initialize the inference engine.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
default_negative_prompt: The default negative prompt to use.
|
||||
"""
|
||||
self.pipeline = pipeline
|
||||
self.inference_args = inference_args
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
@classmethod
|
||||
def create_engine(
|
||||
cls,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> "InferenceEngine":
|
||||
"""
|
||||
Create an inference engine with the specified arguments.
|
||||
|
||||
Args:
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
model_loader_cls: The model loader class to use. If None, it will be
|
||||
determined from the model type.
|
||||
pipeline_type: The type of pipeline to create. If None, it will be
|
||||
@@ -71,16 +71,16 @@ class InferenceEngine:
|
||||
# this way for training we can just do pipeline_cls.from_pretrained(
|
||||
# checkpoint_path) and have it handle everything.
|
||||
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
|
||||
pipeline = build_pipeline(inference_args)
|
||||
pipeline = build_pipeline(fastvideo_args)
|
||||
logger.info("Pipeline Ready")
|
||||
|
||||
# Create the inference engine
|
||||
return cls(pipeline, inference_args)
|
||||
return cls(pipeline, fastvideo_args)
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run inference with the pipeline.
|
||||
@@ -96,17 +96,17 @@ class InferenceEngine:
|
||||
"""
|
||||
out_dict: Dict[str, Any] = dict()
|
||||
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = inference_args.seed
|
||||
height = inference_args.height
|
||||
width = inference_args.width
|
||||
video_length = inference_args.num_frames
|
||||
negative_prompt = inference_args.neg_prompt
|
||||
infer_steps = inference_args.num_inference_steps
|
||||
guidance_scale = inference_args.guidance_scale
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
image_path = inference_args.image_path
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
height = fastvideo_args.height
|
||||
width = fastvideo_args.width
|
||||
video_length = fastvideo_args.num_frames
|
||||
negative_prompt = fastvideo_args.neg_prompt
|
||||
infer_steps = fastvideo_args.num_inference_steps
|
||||
guidance_scale = fastvideo_args.guidance_scale
|
||||
flow_shift = fastvideo_args.flow_shift
|
||||
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
|
||||
image_path = fastvideo_args.image_path
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
@@ -161,21 +161,21 @@ class InferenceEngine:
|
||||
# return
|
||||
# sp_group = get_sp_group()
|
||||
# local_rank = sp_group.rank
|
||||
device = torch.device(inference_args.device_str)
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
image_path=image_path,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
height=inference_args.height,
|
||||
width=inference_args.width,
|
||||
num_frames=inference_args.num_frames,
|
||||
num_inference_steps=inference_args.num_inference_steps,
|
||||
guidance_scale=inference_args.guidance_scale,
|
||||
height=fastvideo_args.height,
|
||||
width=fastvideo_args.width,
|
||||
num_frames=fastvideo_args.num_frames,
|
||||
num_inference_steps=fastvideo_args.num_inference_steps,
|
||||
guidance_scale=fastvideo_args.guidance_scale,
|
||||
# generator=generator,
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if inference_args.num_frames > 1 else "image",
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={}, # Any additional parameters
|
||||
)
|
||||
@@ -184,7 +184,7 @@ class InferenceEngine:
|
||||
print(batch)
|
||||
print('===============================================')
|
||||
print('===============================================')
|
||||
print(inference_args)
|
||||
print(fastvideo_args)
|
||||
|
||||
# ========================================================================
|
||||
# Pipeline inference
|
||||
@@ -192,7 +192,7 @@ class InferenceEngine:
|
||||
start_time = time.time()
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
).output
|
||||
# TODO(will): fix and move to hunyuan stage
|
||||
# out_dict["seeds"] = batch.seeds
|
||||
|
||||
@@ -83,13 +83,13 @@ def get_hf_config(
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
inference_args: Optional[dict] = None,
|
||||
fastvideo_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
model: The model name or path.
|
||||
inference_args: Optional inference arguments to override in the config.
|
||||
fastvideo_args: Optional inference arguments to override in the config.
|
||||
|
||||
Returns:
|
||||
The loaded configuration.
|
||||
|
||||
@@ -13,7 +13,7 @@ from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
|
||||
get_hf_config)
|
||||
@@ -36,14 +36,14 @@ class ComponentLoader(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load the component based on the model path, architecture, and inference args.
|
||||
|
||||
Args:
|
||||
model_path: Path to the component model
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded component
|
||||
@@ -199,20 +199,20 @@ class TextEncoderLoader(ComponentLoader):
|
||||
yield from self._get_weights_iterator(source)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
inference_args.text_encoder_precision)
|
||||
fastvideo_args.text_encoder_precision)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
@@ -249,27 +249,27 @@ class TextEncoderLoader(ComponentLoader):
|
||||
class ImageEncoderLoader(TextEncoderLoader):
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
inference_args.image_encoder_precision)
|
||||
fastvideo_args.image_encoder_precision)
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the image processor based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading image processor from %s", model_path)
|
||||
|
||||
@@ -283,7 +283,7 @@ class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the tokenizer based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading tokenizer from %s", model_path)
|
||||
|
||||
@@ -301,7 +301,7 @@ class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
config = get_diffusers_config(model=model_path)
|
||||
@@ -312,7 +312,7 @@ class VAELoader(ComponentLoader):
|
||||
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(**config).to(inference_args.device)
|
||||
vae = vae_cls(**config).to(fastvideo_args.device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -323,7 +323,7 @@ class VAELoader(ComponentLoader):
|
||||
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(loaded)
|
||||
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae = vae.eval().to(dtype)
|
||||
|
||||
return vae
|
||||
@@ -333,7 +333,7 @@ class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
@@ -354,16 +354,16 @@ class TransformerLoader(ComponentLoader):
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
@@ -380,7 +380,7 @@ class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
@@ -391,8 +391,8 @@ class SchedulerLoader(ComponentLoader):
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
if inference_args.flow_shift is not None:
|
||||
scheduler.set_shift(inference_args.flow_shift)
|
||||
if fastvideo_args.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.flow_shift)
|
||||
|
||||
return scheduler
|
||||
|
||||
@@ -405,7 +405,7 @@ class GenericComponentLoader(ComponentLoader):
|
||||
self.library = library
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load a generic component based on the model path, architecture, and inference args."""
|
||||
logger.warning("Using generic loader for %s with library %s",
|
||||
model_path, self.library)
|
||||
@@ -415,8 +415,8 @@ class GenericComponentLoader(ComponentLoader):
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
)
|
||||
logger.info("Loaded generic transformers model: %s",
|
||||
model.__class__.__name__)
|
||||
@@ -443,7 +443,7 @@ class PipelineComponentLoader:
|
||||
@staticmethod
|
||||
def load_module(module_name: str, component_model_path: str,
|
||||
transformers_or_diffusers: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load a pipeline module.
|
||||
|
||||
@@ -452,7 +452,7 @@ class PipelineComponentLoader:
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
@@ -469,4 +469,4 @@ class PipelineComponentLoader:
|
||||
transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, architecture, inference_args)
|
||||
return loader.load(component_model_path, architecture, fastvideo_args)
|
||||
|
||||
@@ -8,7 +8,7 @@ from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
class BaseScheduler(ABC):
|
||||
timesteps: torch.tensor
|
||||
timesteps: torch.Tensor
|
||||
order: int
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
@@ -38,9 +38,9 @@ class BaseScheduler(ABC):
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
model_output: torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> Union[BaseOutput, Tuple]:
|
||||
pass
|
||||
|
||||
@@ -40,7 +40,7 @@ from fastvideo.v1.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
# Import other required stages
|
||||
)
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
class YourCustomPipeline(ComposedPipelineBase):
|
||||
@@ -53,7 +53,7 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
# Add other required modules
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
# Add and configure pipeline stages
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
@@ -61,14 +61,14 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
)
|
||||
# Add more stages as needed
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Initialize pipeline-specific components
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Implement your pipeline's forward pass
|
||||
batch = self.input_validation_stage(batch, inference_args)
|
||||
batch = self.input_validation_stage(batch, fastvideo_args)
|
||||
# Add more stage executions
|
||||
return batch
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ Diffusion pipelines for fastvideo.v1.
|
||||
This package contains diffusion pipelines for generating videos and images.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -16,7 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
|
||||
"""
|
||||
Only works with valid hf diffusers configs. (model_index.json)
|
||||
We want to build a pipeline based on the inference args mode_path:
|
||||
@@ -25,9 +25,9 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
3. based on the config, determine the pipeline class
|
||||
"""
|
||||
# Get pipeline type
|
||||
model_path = inference_args.model_path
|
||||
model_path = fastvideo_args.model_path
|
||||
model_path = maybe_download_model(model_path)
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
|
||||
@@ -41,7 +41,7 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
pipeline_architecture)
|
||||
|
||||
# instantiate the pipeline
|
||||
pipeline = pipeline_cls(model_path, inference_args, config)
|
||||
pipeline = pipeline_cls(model_path, fastvideo_args, config)
|
||||
logger.info("Pipeline instantiated")
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
|
||||
@@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -38,7 +38,7 @@ class ComposedPipelineBase(ABC):
|
||||
# TODO(will): args should support both inference args and training args
|
||||
def __init__(self,
|
||||
model_path: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None):
|
||||
"""
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
@@ -61,12 +61,12 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
self.modules = self.load_modules(inference_args)
|
||||
self.modules = self.load_modules(fastvideo_args)
|
||||
|
||||
self.initialize_pipeline(inference_args)
|
||||
self.initialize_pipeline(fastvideo_args)
|
||||
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(inference_args)
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
|
||||
def get_module(self, module_name: str) -> Any:
|
||||
return self.modules[module_name]
|
||||
@@ -77,7 +77,7 @@ class ComposedPipelineBase(ABC):
|
||||
def _load_config(self, model_path: str) -> Dict[str, Any]:
|
||||
model_path = maybe_download_model(self.model_path)
|
||||
self.model_path = model_path
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
return cast(Dict[str, Any], config)
|
||||
@@ -108,20 +108,20 @@ class ComposedPipelineBase(ABC):
|
||||
return self._stages
|
||||
|
||||
@abstractmethod
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Create the pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def load_modules(self, inference_args: InferenceArgs) -> Dict[str, Any]:
|
||||
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
"""
|
||||
Load the modules from the config.
|
||||
"""
|
||||
@@ -156,7 +156,7 @@ class ComposedPipelineBase(ABC):
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module %s from %s", module_name,
|
||||
component_model_path)
|
||||
@@ -185,14 +185,14 @@ class ComposedPipelineBase(ABC):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Generate a video or image using the pipeline.
|
||||
|
||||
Args:
|
||||
batch: The batch to generate from.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
Returns:
|
||||
ForwardBatch: The batch with the generated video or image.
|
||||
"""
|
||||
@@ -201,7 +201,7 @@ class ComposedPipelineBase(ABC):
|
||||
self._stage_name_mapping.keys())
|
||||
logger.info("Batch: %s", batch)
|
||||
for stage in self.stages:
|
||||
batch = stage(batch, inference_args)
|
||||
batch = stage(batch, fastvideo_args)
|
||||
|
||||
# Return the output
|
||||
return batch
|
||||
|
||||
@@ -8,7 +8,7 @@ using the modular pipeline architecture.
|
||||
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (CLIPTextEncodingStage,
|
||||
@@ -30,7 +30,7 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
"transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -67,20 +67,20 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
|
||||
1)
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
self.image_processor = VaeImageProcessor(
|
||||
vae_scale_factor=vae_scale_factor)
|
||||
self.add_module("image_processor", self.image_processor)
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = HunyuanVideoPipeline
|
||||
|
||||
@@ -22,7 +22,7 @@ class ForwardBatch:
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
# TODO(will): double check that args are separate from inference_args
|
||||
# TODO(will): double check that args are separate from fastvideo_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
data_type: str
|
||||
|
||||
@@ -12,7 +12,7 @@ from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -45,7 +45,7 @@ class PipelineStage(ABC):
|
||||
def __call__(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
@@ -53,7 +53,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -65,7 +65,7 @@ class PipelineStage(ABC):
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, inference_args)
|
||||
result = self._call_implementation(batch, fastvideo_args)
|
||||
|
||||
execution_time = time.time() - start_time
|
||||
self._logger.info("[%s] Execution completed in %s ms",
|
||||
@@ -85,13 +85,13 @@ class PipelineStage(ABC):
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
# TODO(will): Also handle backward
|
||||
return self.forward(batch, inference_args)
|
||||
return self.forward(batch, fastvideo_args)
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Forward pass of the stage's processing.
|
||||
@@ -101,7 +101,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -111,6 +111,6 @@ class PipelineStage(ABC):
|
||||
def backward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -8,7 +8,7 @@ This module contains implementations of image encoding stages for diffusion pipe
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -40,19 +40,19 @@ class CLIPImageEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into image encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(batch.device)
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
@@ -64,7 +64,7 @@ class CLIPImageEncodingStage(PipelineStage):
|
||||
|
||||
batch.image_embeds.append(image_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -39,19 +39,19 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -85,7 +85,7 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ Conditioning stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -24,14 +24,14 @@ class ConditioningStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Apply conditioning to the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with applied conditioning.
|
||||
|
||||
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -28,14 +28,14 @@ class DecodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
@@ -46,13 +46,13 @@ class DecodingStage(PipelineStage):
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if inference_args.output_type == "latent":
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
@@ -73,9 +73,9 @@ class DecodingStage(PipelineStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
|
||||
@@ -17,7 +17,7 @@ from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather)
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -51,20 +51,20 @@ class DenoisingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# If use cpu offload, need to load the model back into gpu again
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer = self.transformer.to(batch.device)
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
@@ -76,9 +76,9 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Handle sequence parallelism if enabled
|
||||
world_size, rank = get_sequence_model_parallel_world_size(
|
||||
@@ -161,11 +161,11 @@ class DenoisingStage(PipelineStage):
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[inference_args.embedded_cfg_scale] *
|
||||
[fastvideo_args.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=batch.device,
|
||||
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale
|
||||
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
@@ -193,7 +193,7 @@ class DenoisingStage(PipelineStage):
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
current_timestep=i,
|
||||
forward_batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
@@ -203,11 +203,11 @@ class DenoisingStage(PipelineStage):
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# inference_args or training_args, and attn_metadata.
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = self.transformer(
|
||||
@@ -223,7 +223,7 @@ class DenoisingStage(PipelineStage):
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred_uncond = self.transformer(
|
||||
@@ -267,7 +267,7 @@ class DenoisingStage(PipelineStage):
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Optional
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
@@ -33,14 +33,14 @@ class EncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode pixel representations into latent space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
@@ -62,7 +62,7 @@ class EncodingStage(PipelineStage):
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
inference_args.num_frames - 1, batch.height,
|
||||
fastvideo_args.num_frames - 1, batch.height,
|
||||
batch.width)
|
||||
],
|
||||
dim=2)
|
||||
@@ -70,17 +70,17 @@ class EncodingStage(PipelineStage):
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not inference_args.disable_autocast
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
@@ -106,9 +106,9 @@ class EncodingStage(PipelineStage):
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, inference_args.num_frames,
|
||||
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size[:, :, list(range(1, inference_args.num_frames))] = 0
|
||||
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
|
||||
@@ -5,7 +5,7 @@ Input validation stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -22,10 +22,10 @@ class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = inference_args.seed
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
@@ -37,19 +37,19 @@ class InputValidationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Validate and prepare inputs.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The validated batch information.
|
||||
"""
|
||||
self._generate_seeds(batch, inference_args)
|
||||
self._generate_seeds(batch, fastvideo_args)
|
||||
|
||||
# Ensure prompt is properly formatted
|
||||
if batch.prompt is None and batch.prompt_embeds is None:
|
||||
@@ -91,6 +91,6 @@ class InputValidationStage(PipelineStage):
|
||||
|
||||
# Set data type if not already set
|
||||
if batch.data_type is None:
|
||||
batch.data_type = inference_args.precision
|
||||
batch.data_type = fastvideo_args.precision
|
||||
|
||||
return batch
|
||||
|
||||
@@ -4,7 +4,7 @@ Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -29,14 +29,14 @@ class LatentPreparationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare initial latent variables for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared latent variables.
|
||||
@@ -44,7 +44,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Adjust video length based on VAE version if needed
|
||||
if hasattr(self, 'adjust_video_length'):
|
||||
batch = self.adjust_video_length(self.vae, batch, inference_args)
|
||||
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -69,16 +69,16 @@ class LatentPreparationStage(PipelineStage):
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
assert inference_args.num_channels_latents is not None
|
||||
assert inference_args.vae_scale_factor is not None
|
||||
assert fastvideo_args.num_channels_latents is not None
|
||||
assert fastvideo_args.vae_scale_factor is not None
|
||||
|
||||
# Calculate latent shape
|
||||
shape = (
|
||||
batch_size,
|
||||
inference_args.num_channels_latents,
|
||||
fastvideo_args.num_channels_latents,
|
||||
num_frames,
|
||||
height // inference_args.vae_scale_factor,
|
||||
width // inference_args.vae_scale_factor,
|
||||
height // fastvideo_args.vae_scale_factor,
|
||||
width // fastvideo_args.vae_scale_factor,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
@@ -107,13 +107,13 @@ class LatentPreparationStage(PipelineStage):
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs) -> ForwardBatch:
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with adjusted video length.
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import TypedDict
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -61,19 +61,19 @@ class LlamaEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = prompt_template_video["template"].format(batch.prompt)
|
||||
@@ -123,7 +123,7 @@ class LlamaEncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_last_hidden_state)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -38,19 +38,19 @@ class T5EncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = batch.prompt
|
||||
@@ -109,7 +109,7 @@ class T5EncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(neg_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -29,14 +29,14 @@ class TimestepPreparationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare timesteps for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared timesteps.
|
||||
|
||||
@@ -6,7 +6,7 @@ This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
@@ -24,7 +24,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -65,15 +65,15 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").out_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
|
||||
@@ -6,7 +6,7 @@ This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
@@ -26,7 +26,7 @@ class WanPipeline(ComposedPipelineBase):
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -58,15 +58,15 @@ class WanPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
|
||||
@@ -11,12 +11,12 @@ from einops import rearrange
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.inference_args import InferenceArgs, prepare_inference_args
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
|
||||
# Fix the import path
|
||||
from fastvideo.v1.inference_engine import InferenceEngine
|
||||
|
||||
|
||||
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
@@ -25,33 +25,33 @@ def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
inference_args.device_str = device_str
|
||||
inference_args.device = torch.device(device_str)
|
||||
assert inference_args.sp_size is not None
|
||||
assert inference_args.tp_size is not None
|
||||
fastvideo_args.device_str = device_str
|
||||
fastvideo_args.device = torch.device(device_str)
|
||||
assert fastvideo_args.sp_size is not None
|
||||
assert fastvideo_args.tp_size is not None
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=inference_args.sp_size,
|
||||
tensor_model_parallel_size=inference_args.tp_size,
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size,
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
)
|
||||
|
||||
|
||||
def main(inference_args: InferenceArgs):
|
||||
initialize_distributed_and_parallelism(inference_args)
|
||||
engine = InferenceEngine.create_engine(inference_args, )
|
||||
def main(fastvideo_args: FastVideoArgs):
|
||||
initialize_distributed_and_parallelism(fastvideo_args)
|
||||
engine = InferenceEngine.create_engine(fastvideo_args, )
|
||||
|
||||
if inference_args.prompt_path is not None:
|
||||
with open(inference_args.prompt_path) as f:
|
||||
if fastvideo_args.prompt_path is not None:
|
||||
with open(fastvideo_args.prompt_path) as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
else:
|
||||
if inference_args.prompt is None:
|
||||
if fastvideo_args.prompt is None:
|
||||
raise ValueError("prompt or prompt_path is required")
|
||||
prompts = [inference_args.prompt]
|
||||
prompts = [fastvideo_args.prompt]
|
||||
|
||||
# Process each prompt
|
||||
for prompt in prompts:
|
||||
outputs = engine.run(
|
||||
prompt=prompt,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
# Process outputs
|
||||
@@ -63,13 +63,13 @@ def main(inference_args: InferenceArgs):
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(inference_args.output_path,
|
||||
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(fastvideo_args.output_path,
|
||||
f"{prompt[:100]}.mp4"),
|
||||
frames,
|
||||
fps=inference_args.fps)
|
||||
fps=fastvideo_args.fps)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
inference_args = prepare_inference_args(sys.argv[1:])
|
||||
main(inference_args)
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
main(fastvideo_args)
|
||||
|
||||
@@ -11,7 +11,7 @@ from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
@@ -38,7 +38,7 @@ def test_clip_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="openai/clip-vit-large-patch14",
|
||||
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
|
||||
precision="float16")
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from transformers import AutoConfig
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -38,7 +38,7 @@ def test_llama_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
precision="float16")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch
|
||||
from diffusers import WanTransformer3DModel
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -29,7 +29,7 @@ def test_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=TRANSFORMER_PATH,
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
use_cpu_offload=False,
|
||||
precision=precision_str)
|
||||
args.device = device
|
||||
|
||||
@@ -6,7 +6,7 @@ import pytest
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -28,7 +28,7 @@ def test_wan_vae():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
args.device = device
|
||||
|
||||
loader = VAELoader()
|
||||
|
||||
+104
-7
@@ -10,8 +10,10 @@ import math
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from functools import wraps
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast
|
||||
from functools import wraps, partial
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast, Callable
|
||||
from dataclasses import asdict, fields
|
||||
import cloudpickle
|
||||
|
||||
import filelock
|
||||
import torch
|
||||
@@ -25,7 +27,7 @@ logger = init_logger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
# TODO(will): used to convert inference_args.precision to torch.dtype. Find a
|
||||
# TODO(will): used to convert fastvideo_args.precision to torch.dtype. Find a
|
||||
# cleaner way to do this.
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
@@ -364,9 +366,7 @@ def import_pynvml():
|
||||
install FastVideo. It provides a Python module named `pynvml`.
|
||||
- `pynvml` (https://pypi.org/project/pynvml/): An unofficial wrapper.
|
||||
Prior to version 12.0, it also provides a Python module `pynvml`,
|
||||
and therefore conflicts with the official one. What's worse,
|
||||
the module is a Python package, and has higher priority than
|
||||
the official one which is a standalone Python file.
|
||||
and therefore conflicts with the official one which is a standalone Python file.
|
||||
This causes errors when both of them are installed.
|
||||
Starting from version 12.0, it migrates to a new module
|
||||
named `pynvml_utils` to avoid the conflict.
|
||||
@@ -383,12 +383,15 @@ def import_pynvml():
|
||||
|
||||
|
||||
def maybe_download_model(model_path: str,
|
||||
local_dir: Optional[str] = None) -> str:
|
||||
local_dir: Optional[str] = None,
|
||||
download: bool = True) -> str:
|
||||
"""
|
||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
||||
|
||||
Args:
|
||||
model_path: Local path or Hugging Face Hub model ID
|
||||
local_dir: Local directory to save the model
|
||||
download: Whether to download the model from Hugging Face Hub
|
||||
|
||||
Returns:
|
||||
Local path to the model
|
||||
@@ -457,3 +460,97 @@ def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
|
||||
|
||||
logger.info("Diffusers version: %s", config["_diffusers_version"])
|
||||
return cast(Dict[str, Any], config)
|
||||
|
||||
|
||||
def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Download and extract just the model_index.json for a Hugging Face model.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Path or HF Hub model ID
|
||||
|
||||
Returns:
|
||||
The parsed model_index.json as a dictionary
|
||||
"""
|
||||
import tempfile
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
# If it's a local path, verify it directly
|
||||
if os.path.exists(model_name_or_path):
|
||||
return verify_model_config_and_directory(model_name_or_path)
|
||||
|
||||
# For remote models, download just the model_index.json
|
||||
try:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
# Download just the model_index.json file
|
||||
model_index_path = hf_hub_download(repo_id=model_name_or_path,
|
||||
filename="model_index.json",
|
||||
local_dir=tmp_dir)
|
||||
|
||||
# Load the model_index.json
|
||||
with open(model_index_path) as f:
|
||||
config: Dict[str, Any] = json.load(f)
|
||||
|
||||
# Verify it has the required fields
|
||||
if "_class_name" not in config:
|
||||
raise ValueError(
|
||||
f"model_index.json for {model_name_or_path} does not contain _class_name field"
|
||||
)
|
||||
|
||||
if "_diffusers_version" not in config:
|
||||
raise ValueError(
|
||||
f"model_index.json for {model_name_or_path} does not contain _diffusers_version field"
|
||||
)
|
||||
|
||||
# Add the pipeline name for downstream use
|
||||
config["pipeline_name"] = config["_class_name"]
|
||||
|
||||
logger.info("Downloaded model_index.json for %s, pipeline: %s",
|
||||
model_name_or_path, config["_class_name"])
|
||||
return config
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Failed to download or parse model_index.json for {model_name_or_path}: {e}"
|
||||
) from e
|
||||
|
||||
|
||||
def update_environment_variables(envs: Dict[str, str]):
|
||||
for k, v in envs.items():
|
||||
if k in os.environ and os.environ[k] != v:
|
||||
logger.warning(
|
||||
"Overwriting environment variable %s "
|
||||
"from '%s' to '%s'", k, os.environ[k], v)
|
||||
os.environ[k] = v
|
||||
|
||||
|
||||
def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
|
||||
kwargs: dict[str, Any]) -> Any:
|
||||
"""
|
||||
Run a method of an object with the given arguments and keyword arguments.
|
||||
If the method is string, it will be converted to a method using getattr.
|
||||
If the method is serialized bytes and will be deserialized using
|
||||
cloudpickle.
|
||||
If the method is a callable, it will be called directly.
|
||||
"""
|
||||
if isinstance(method, bytes):
|
||||
func = partial(cloudpickle.loads(method), obj)
|
||||
elif isinstance(method, str):
|
||||
try:
|
||||
func = getattr(obj, method)
|
||||
except AttributeError:
|
||||
raise NotImplementedError(f"Method {method!r} is not"
|
||||
" implemented.") from None
|
||||
else:
|
||||
func = partial(method, obj) # type: ignore
|
||||
return func(*args, **kwargs)
|
||||
|
||||
|
||||
def diff_keys(a, b):
|
||||
return [k for k in asdict(a) if asdict(a)[k] != asdict(b)[k]]
|
||||
|
||||
|
||||
def update_in_place(target, source, ignore_fields=()):
|
||||
for f in fields(target):
|
||||
if hasattr(source, f.name) and f.name not in list(ignore_fields):
|
||||
setattr(target, f.name, getattr(source, f.name))
|
||||
|
||||
+2
-2
@@ -36,8 +36,8 @@ dependencies = [
|
||||
"wandb==0.18.5", "loguru", "test-tube==0.7.5",
|
||||
|
||||
# Miscellaneous Utilities
|
||||
"tqdm==4.66.5", "PyYAML==6.0.1", "idna==3.6", "protobuf==5.28.3",
|
||||
"gradio==5.3.0", "moviepy==1.0.3", "flask",
|
||||
"tqdm==4.66.5", "PyYAML==6.0.1", "protobuf==5.28.3",
|
||||
"gradio>=5.22.0", "moviepy==1.0.3", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch",
|
||||
|
||||
Reference in New Issue
Block a user