Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3def003890 |
@@ -1,40 +0,0 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
use_cpu_offload=True,
|
||||
text_encoder_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
sampling_param.num_frames = 3 # Reduce from default 81 to 25 frames bc we have to use the SDPA attn backend for mps
|
||||
sampling_param.height = 256
|
||||
sampling_param.width = 256
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,39 +0,0 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig, SamplingParam
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_fp16"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
pipeline_config = PipelineConfig.from_pretrained(model)
|
||||
pipeline_config.text_encoder_precisions = ("bf16", )
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model,
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
pipeline_config=pipeline_config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
use_cpu_offload=True,
|
||||
text_encoder_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model)
|
||||
sampling_param.num_frames = 30
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,148 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/main/vllm/distributed/device_communicators/cpu_communicator.py
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.v1.platforms.interface import CpuArchEnum
|
||||
|
||||
from .base_device_communicator import DeviceCommunicatorBase
|
||||
|
||||
|
||||
class CpuCommunicator(DeviceCommunicatorBase):
|
||||
|
||||
def __init__(self,
|
||||
cpu_group: ProcessGroup,
|
||||
device: torch.device | None = None,
|
||||
device_group: ProcessGroup | None = None,
|
||||
unique_name: str = ""):
|
||||
super().__init__(cpu_group, device, device_group, unique_name)
|
||||
self.dist_module = torch.distributed
|
||||
|
||||
if (current_platform.get_cpu_architecture()
|
||||
== CpuArchEnum.X86) and hasattr(
|
||||
torch.ops._C,
|
||||
"init_shm_manager") and unique_name.startswith("tp"):
|
||||
self.dist_module = _CPUSHMDistributed(self)
|
||||
|
||||
def all_reduce(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: torch.distributed.ReduceOp | None = torch.distributed.ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
self.dist_module.all_reduce(input_, group=self.device_group, op=op)
|
||||
return input_
|
||||
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> torch.Tensor | None:
|
||||
"""
|
||||
NOTE: We assume that the input tensor is on the same device across
|
||||
all the ranks.
|
||||
NOTE: `dst` is the local rank of the destination rank.
|
||||
"""
|
||||
world_size = self.world_size
|
||||
assert -input_.dim() <= dim < input_.dim(), (
|
||||
f"Invalid dim ({dim}) for input tensor with shape {input_.size()}")
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
|
||||
# Allocate output tensor.
|
||||
if self.rank_in_group == dst:
|
||||
gather_list = [torch.empty_like(input_) for _ in range(world_size)]
|
||||
else:
|
||||
gather_list = None
|
||||
|
||||
# Gather.
|
||||
self.dist_module.gather(input_,
|
||||
gather_list,
|
||||
dst=self.ranks[dst],
|
||||
group=self.device_group)
|
||||
|
||||
if self.rank_in_group == dst:
|
||||
output_tensor = torch.cat(gather_list, dim=dim)
|
||||
else:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
input_size = input_.size()
|
||||
# NOTE: we have to use concat-style all-gather here,
|
||||
# stack-style all-gather has compatibility issues with
|
||||
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
|
||||
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
# All-gather.
|
||||
self.dist_module.all_gather_into_tensor(output_tensor,
|
||||
input_,
|
||||
group=self.device_group)
|
||||
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(self.world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
|
||||
|
||||
class _CPUSHMDistributed:
|
||||
|
||||
def __init__(self, communicator: CpuCommunicator):
|
||||
instance_identifier = os.environ["VLLM_DIST_IDENT"]
|
||||
unique_name = communicator.unique_name
|
||||
instance_identifier = f"{instance_identifier}-{unique_name}"
|
||||
self.communicator = communicator
|
||||
|
||||
group_ranks = [str(rank) for rank in self.communicator.ranks]
|
||||
shm_group_identifier = f"[{'-'.join(group_ranks)}]"
|
||||
self.group_name = f"{instance_identifier}-{shm_group_identifier}-cpushm"
|
||||
|
||||
self.handle = self._init_cpu_shm()
|
||||
|
||||
def _init_cpu_shm(self) -> int:
|
||||
handle = torch.ops._C.init_shm_manager(
|
||||
self.group_name,
|
||||
self.communicator.world_size,
|
||||
self.communicator.rank,
|
||||
)
|
||||
torch.distributed.barrier(self.communicator.device_group)
|
||||
torch.ops._C.join_shm_manager(
|
||||
handle,
|
||||
self.group_name,
|
||||
)
|
||||
torch.distributed.barrier(self.communicator.device_group)
|
||||
|
||||
return int(handle)
|
||||
|
||||
def all_reduce(self,
|
||||
input: torch.Tensor,
|
||||
group: ProcessGroup | None = None) -> None:
|
||||
torch.ops._C.shm_allreduce(self.handle, input)
|
||||
|
||||
def gather(self,
|
||||
input: torch.Tensor,
|
||||
gather_list: list[torch.Tensor] | None,
|
||||
dst: int = -1,
|
||||
group: ProcessGroup | None = None) -> None:
|
||||
# Note: different from the torch gather, here we use local dst rank.
|
||||
torch.ops._C.shm_gather(self.handle, input, gather_list,
|
||||
torch.distributed.get_group_rank(group, dst))
|
||||
|
||||
def all_gather_into_tensor(self,
|
||||
output: torch.Tensor,
|
||||
input: torch.Tensor,
|
||||
group: ProcessGroup | None = None) -> None:
|
||||
torch.ops._C.shm_all_gather(self.handle, input, output)
|
||||
@@ -102,8 +102,7 @@ class PyNcclCommunicator:
|
||||
# A small all_reduce for warmup.
|
||||
data = torch.zeros(1, device=device)
|
||||
self.all_reduce(data)
|
||||
if stream is not None:
|
||||
stream.synchronize()
|
||||
stream.synchronize()
|
||||
del data
|
||||
|
||||
def all_reduce(self,
|
||||
|
||||
@@ -41,18 +41,17 @@ from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase)
|
||||
from fastvideo.v1.distributed.device_communicators.cpu_communicator import (
|
||||
CpuCommunicator)
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
from fastvideo.v1.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GraphCaptureContext:
|
||||
stream: torch.cuda.Stream | None
|
||||
stream: torch.cuda.Stream
|
||||
|
||||
|
||||
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
|
||||
@@ -184,30 +183,23 @@ class GroupCoordinator:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
# TODO: fix it for other platforms
|
||||
self.device = get_local_torch_device()
|
||||
if current_platform.is_cuda_alike():
|
||||
self.device = torch.device(f"cuda:{local_rank}")
|
||||
else:
|
||||
self.device = torch.device("cpu")
|
||||
|
||||
self.use_device_communicator = use_device_communicator
|
||||
|
||||
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
||||
if use_device_communicator and self.world_size > 1:
|
||||
# Platform-aware device communicator selection
|
||||
if current_platform.is_cuda_alike():
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
self.device_communicator = CudaCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
else:
|
||||
# For MPS and CPU, use the CPU communicator
|
||||
self.device_communicator = CpuCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
# device_comm_cls = resolve_obj_by_qualname(
|
||||
# current_platform.get_device_communicator_cls())
|
||||
self.device_communicator = CudaCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
|
||||
self.mq_broadcaster = None
|
||||
|
||||
@@ -254,29 +246,19 @@ class GroupCoordinator:
|
||||
@contextmanager
|
||||
def graph_capture(self,
|
||||
graph_capture_context: GraphCaptureContext | None = None):
|
||||
# Platform-aware graph capture
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
if graph_capture_context is None:
|
||||
stream = torch.cuda.Stream()
|
||||
graph_capture_context = GraphCaptureContext(stream)
|
||||
else:
|
||||
stream = graph_capture_context.stream
|
||||
|
||||
# ensure all initialization operations complete before attempting to
|
||||
# capture the graph on another stream
|
||||
curr_stream = torch.cuda.current_stream()
|
||||
if curr_stream != stream:
|
||||
stream.wait_stream(curr_stream)
|
||||
|
||||
with torch.cuda.stream(stream):
|
||||
yield graph_capture_context
|
||||
if graph_capture_context is None:
|
||||
stream = torch.cuda.Stream()
|
||||
graph_capture_context = GraphCaptureContext(stream)
|
||||
else:
|
||||
# For non-CUDA platforms (MPS, CPU), just yield the context without stream management
|
||||
if graph_capture_context is None:
|
||||
# Create a dummy context for non-CUDA platforms
|
||||
graph_capture_context = GraphCaptureContext(None)
|
||||
stream = graph_capture_context.stream
|
||||
|
||||
# ensure all initialization operations complete before attempting to
|
||||
# capture the graph on another stream
|
||||
curr_stream = torch.cuda.current_stream()
|
||||
if curr_stream != stream:
|
||||
stream.wait_stream(curr_stream)
|
||||
|
||||
with torch.cuda.stream(stream):
|
||||
yield graph_capture_context
|
||||
|
||||
def all_reduce(
|
||||
@@ -710,7 +692,6 @@ class GroupCoordinator:
|
||||
|
||||
|
||||
_WORLD: GroupCoordinator | None = None
|
||||
_NODE: GroupCoordinator | None = None
|
||||
|
||||
|
||||
def get_world_group() -> GroupCoordinator:
|
||||
@@ -771,14 +752,6 @@ def init_distributed_environment(
|
||||
backend: str = "nccl",
|
||||
device_id: torch.device | None = None,
|
||||
):
|
||||
# Determine the appropriate backend based on the platform
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
if backend == "nccl" and not current_platform.is_cuda_alike():
|
||||
# Use gloo backend for non-CUDA platforms (MPS, CPU)
|
||||
backend = "gloo"
|
||||
logger.info("Using gloo backend for %s platform",
|
||||
current_platform.device_name)
|
||||
|
||||
logger.debug(
|
||||
"world_size=%d rank=%d local_rank=%d "
|
||||
"distributed_init_method=%s backend=%s", world_size, rank, local_rank,
|
||||
@@ -787,22 +760,13 @@ def init_distributed_environment(
|
||||
assert distributed_init_method is not None, (
|
||||
"distributed_init_method must be provided when initializing "
|
||||
"distributed environment")
|
||||
|
||||
# For MPS, don't pass device_id as it doesn't support device indices
|
||||
if current_platform.is_mps():
|
||||
torch.distributed.init_process_group(
|
||||
backend=backend,
|
||||
init_method=distributed_init_method,
|
||||
world_size=world_size,
|
||||
rank=rank)
|
||||
else:
|
||||
# this backend is used for WORLD
|
||||
torch.distributed.init_process_group(
|
||||
backend=backend,
|
||||
init_method=distributed_init_method,
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
device_id=device_id)
|
||||
# this backend is used for WORLD
|
||||
torch.distributed.init_process_group(
|
||||
backend=backend,
|
||||
init_method=distributed_init_method,
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
device_id=device_id)
|
||||
# set the local rank
|
||||
# local_rank is not available in torch ProcessGroup,
|
||||
# see https://github.com/pytorch/pytorch/issues/122816
|
||||
@@ -945,9 +909,7 @@ def get_dp_rank() -> int:
|
||||
|
||||
def get_local_torch_device() -> torch.device:
|
||||
"""Return the torch device for the current rank."""
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}"
|
||||
) if current_platform.is_cuda_alike() else torch.device(
|
||||
"mps")
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
|
||||
def maybe_init_distributed_environment_and_model_parallel(
|
||||
@@ -962,8 +924,7 @@ def maybe_init_distributed_environment_and_model_parallel(
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
device = get_local_torch_device()
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
@@ -973,11 +934,7 @@ def maybe_init_distributed_environment_and_model_parallel(
|
||||
device_id=device)
|
||||
initialize_model_parallel(tensor_model_parallel_size=tp_size,
|
||||
sequence_model_parallel_size=sp_size)
|
||||
|
||||
# Only set CUDA device if we're on a CUDA platform
|
||||
if current_platform.is_cuda_alike():
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
torch.cuda.set_device(device)
|
||||
|
||||
|
||||
def model_parallel_is_initialized() -> bool:
|
||||
@@ -1060,8 +1017,8 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
ray.shutdown()
|
||||
|
||||
|
||||
def is_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
|
||||
source_rank: int = 0) -> list[int]:
|
||||
def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
|
||||
source_rank: int = 0) -> list[bool]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
as the source rank. It tests if processes are attached to the same
|
||||
|
||||
@@ -10,7 +10,6 @@ from typing import Any
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -80,13 +79,6 @@ class FastVideoArgs:
|
||||
# Stage verification
|
||||
enable_stage_verification: bool = True
|
||||
|
||||
# model paths for correct deallocation
|
||||
model_paths: dict[str, str] = field(default_factory=dict)
|
||||
model_loaded: dict[str, bool] = field(default_factory=lambda: {
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
@@ -281,15 +273,7 @@ class FastVideoArgs:
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
# Get the field to check if it has a default_factory
|
||||
field = dataclasses.fields(cls)[next(
|
||||
i for i, f in enumerate(dataclasses.fields(cls))
|
||||
if f.name == attr)]
|
||||
if field.default_factory is not dataclasses.MISSING:
|
||||
# Use the default_factory to create the default value
|
||||
default_value = field.default_factory()
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
@@ -302,9 +286,6 @@ class FastVideoArgs:
|
||||
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if current_platform.is_mps():
|
||||
self.use_fsdp_inference = False
|
||||
|
||||
if not self.inference_mode:
|
||||
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
|
||||
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
|
||||
@@ -482,31 +463,16 @@ class TrainingArgs(FastVideoArgs):
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
logger.info(provided_args)
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs: dict[str, Any] = {}
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
if attr == 'pipeline_config':
|
||||
pipeline_config = PipelineConfig.from_kwargs(provided_args)
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
# Get the field to check its default value
|
||||
field = dataclasses.fields(cls)[next(
|
||||
i for i, f in enumerate(dataclasses.fields(cls))
|
||||
if f.name == attr)]
|
||||
|
||||
# Check if the attribute is provided in args
|
||||
if hasattr(args, attr):
|
||||
value = getattr(args, attr)
|
||||
else:
|
||||
# Use the field's default value
|
||||
if field.default_factory is not dataclasses.MISSING:
|
||||
value = field.default_factory()
|
||||
elif field.default is not dataclasses.MISSING:
|
||||
value = field.default
|
||||
else:
|
||||
# No default value, use None
|
||||
value = None
|
||||
|
||||
kwargs[attr] = value
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
|
||||
@CustomOp.register("rms_norm")
|
||||
@@ -34,8 +33,7 @@ class RMSNorm(CustomOp):
|
||||
else var_hidden_size)
|
||||
self.has_weight = has_weight
|
||||
|
||||
self.weight = torch.ones(hidden_size) if current_platform.is_cuda_alike(
|
||||
) else torch.ones(hidden_size, dtype=dtype)
|
||||
self.weight = torch.ones(hidden_size)
|
||||
if self.has_weight:
|
||||
self.weight = nn.Parameter(self.weight)
|
||||
|
||||
|
||||
@@ -103,11 +103,9 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
weight = Parameter(torch.empty(sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype),
|
||||
requires_grad=False)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
@@ -117,11 +115,8 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
output = F.linear(x, layer.weight, bias) if torch.cuda.is_available(
|
||||
) or bias is None else F.linear(
|
||||
x, layer.weight, bias.to(x.dtype)
|
||||
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
|
||||
return output
|
||||
|
||||
return F.linear(x, layer.weight, bias)
|
||||
|
||||
|
||||
class LinearBase(torch.nn.Module):
|
||||
@@ -205,10 +200,7 @@ class ReplicatedLinear(LinearBase):
|
||||
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(
|
||||
self.output_size,
|
||||
dtype=self.params_dtype,
|
||||
))
|
||||
torch.empty(self.output_size, dtype=self.params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
@@ -309,10 +301,7 @@ class ColumnParallelLinear(LinearBase):
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(
|
||||
self.output_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
))
|
||||
torch.empty(self.output_size_per_partition, dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
@@ -516,8 +505,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
if isinstance(param, (PackedColumnParameter | PackedvLLMParameter
|
||||
)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
@@ -683,8 +672,8 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
if isinstance(param, (PackedColumnParameter | PackedvLLMParameter
|
||||
)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
|
||||
@@ -300,4 +300,4 @@ def replace_submodule(model: nn.Module, module_name: str,
|
||||
parent = model.get_submodule(".".join(module_name.split(".")[:-1]))
|
||||
target_name = module_name.split(".")[-1]
|
||||
setattr(parent, target_name, new_module)
|
||||
return new_module
|
||||
return new_module
|
||||
|
||||
@@ -324,7 +324,7 @@ def get_nd_rotary_pos_embed(
|
||||
else:
|
||||
grid = full_grid
|
||||
|
||||
if isinstance(theta_rescale_factor, int | float):
|
||||
if isinstance(theta_rescale_factor, (int | float)):
|
||||
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||
elif isinstance(theta_rescale_factor,
|
||||
list) and len(theta_rescale_factor) == 1:
|
||||
@@ -333,7 +333,7 @@ def get_nd_rotary_pos_embed(
|
||||
rope_dim_list
|
||||
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
if isinstance(interpolation_factor, int | float):
|
||||
if isinstance(interpolation_factor, (int | float)):
|
||||
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||
elif isinstance(interpolation_factor,
|
||||
list) and len(interpolation_factor) == 1:
|
||||
|
||||
@@ -35,7 +35,7 @@ class PatchEmbed(nn.Module):
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if isinstance(patch_size, (list | tuple)):
|
||||
if len(patch_size) == 1:
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
|
||||
@@ -27,12 +27,9 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs):
|
||||
"""Create weights for embedding layer."""
|
||||
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
weight = Parameter(torch.empty(sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype),
|
||||
requires_grad=False)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
@@ -54,7 +54,7 @@ class PatchEmbed2D(nn.Module):
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if isinstance(patch_size, (list | tuple)):
|
||||
if len(patch_size) == 1:
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
|
||||
@@ -25,11 +25,8 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
@@ -627,7 +624,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
|
||||
dtype=torch.float64,
|
||||
rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
@@ -645,10 +642,6 @@ class WanTransformer3DModel(CachableDiT):
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states.to(
|
||||
orig_dtype) if current_platform.is_mps(
|
||||
) else encoder_hidden_states # cast to orig_dtype for MPS
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
# 4. Transformer blocks
|
||||
|
||||
@@ -37,7 +37,6 @@ from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.encoders.base import TextEncoder
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
|
||||
class AttentionType:
|
||||
@@ -326,9 +325,8 @@ class T5Attention(nn.Module):
|
||||
attention_mask = attention_mask.view(
|
||||
bs, 1, 1,
|
||||
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
|
||||
mask_val = -1e4 if current_platform.is_mps() else torch.finfo(
|
||||
q.dtype).min
|
||||
attn_bias.masked_fill_(attention_mask == 0, mask_val)
|
||||
attn_bias.masked_fill_(attention_mask == 0,
|
||||
torch.finfo(q.dtype).min)
|
||||
|
||||
if get_tp_world_size() > 1:
|
||||
rank = get_tp_rank()
|
||||
|
||||
@@ -30,7 +30,6 @@ from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
pt_weights_iterator, safetensors_weights_iterator)
|
||||
from fastvideo.v1.models.registry import ModelRegistry
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -252,8 +251,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
|
||||
|
||||
if fastvideo_args.text_encoder_offload:
|
||||
target_device = torch.device(
|
||||
"mps") if current_platform.is_mps() else torch.device("cpu")
|
||||
target_device = torch.device("cpu")
|
||||
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
@@ -271,28 +269,18 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self.counter_after_loading_weights -
|
||||
self.counter_before_loading_weights)
|
||||
|
||||
# Explicitly move model to target device after loading weights
|
||||
model = model.to(target_device)
|
||||
|
||||
if use_cpu_offload:
|
||||
# Disable FSDP for MPS as it's not compatible
|
||||
if current_platform.is_mps():
|
||||
logger.info(
|
||||
"Disabling FSDP sharding for MPS platform as it's not compatible"
|
||||
)
|
||||
else:
|
||||
mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(1, dist.get_world_size()),
|
||||
mesh_dim_names=("offload", "replicate"),
|
||||
)
|
||||
shard_model(
|
||||
model,
|
||||
cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
mesh=mesh["offload"],
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
|
||||
mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(1, dist.get_world_size()),
|
||||
mesh_dim_names=("offload", "replicate"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
mesh=mesh["offload"],
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
|
||||
# We only enable strict check for non-quantized models
|
||||
# that have loaded weights tracking currently.
|
||||
# if loaded_weights is not None:
|
||||
@@ -371,7 +359,6 @@ class VAELoader(ComponentLoader):
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
fastvideo_args.model_paths["vae"] = model_path
|
||||
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
@@ -408,8 +395,6 @@ class TransformerLoader(ComponentLoader):
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
|
||||
fastvideo_args.model_paths["transformer"] = model_path
|
||||
|
||||
# Config from Diffusers supersedes fastvideo's model config
|
||||
dit_config = fastvideo_args.pipeline_config.dit_config
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
@@ -89,35 +89,23 @@ def maybe_load_fsdp_model(
|
||||
|
||||
with set_default_dtype(param_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
# Check if we should use FSDP
|
||||
use_fsdp = training_mode or fsdp_inference
|
||||
|
||||
# Disable FSDP for MPS as it's not compatible
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
if current_platform.is_mps():
|
||||
use_fsdp = False
|
||||
logger.info("Disabling FSDP for MPS platform as it's not compatible")
|
||||
|
||||
if use_fsdp:
|
||||
world_size = hsdp_replicate_dim * hsdp_shard_dim
|
||||
if not training_mode and not fsdp_inference:
|
||||
hsdp_replicate_dim = world_size
|
||||
hsdp_shard_dim = 1
|
||||
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
# (Replicate(), Shard(dim=0))
|
||||
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh,
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=pin_cpu_memory)
|
||||
world_size = hsdp_replicate_dim * hsdp_shard_dim
|
||||
if not training_mode and not fsdp_inference:
|
||||
hsdp_replicate_dim = world_size
|
||||
hsdp_shard_dim = 1
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
# (Replicate(), Shard(dim=0))
|
||||
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh,
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=pin_cpu_memory)
|
||||
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
|
||||
|
||||
@@ -113,8 +113,9 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
if shard_offset is None or shard_size is None:
|
||||
raise ValueError("shard_offset and shard_size must be provided")
|
||||
if isinstance(
|
||||
self, PackedColumnParameter
|
||||
| PackedvLLMParameter) and self.packed_dim == self.output_dim:
|
||||
self,
|
||||
(PackedColumnParameter
|
||||
| PackedvLLMParameter)) and self.packed_dim == self.output_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
@@ -141,8 +142,9 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
assert num_heads is not None
|
||||
|
||||
if isinstance(
|
||||
self, PackedColumnParameter
|
||||
| PackedvLLMParameter) and self.output_dim == self.packed_dim:
|
||||
self,
|
||||
(PackedColumnParameter
|
||||
| PackedvLLMParameter)) and self.output_dim == self.packed_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
|
||||
@@ -259,7 +259,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
assert self._step_index is not None
|
||||
self._step_index += 1
|
||||
|
||||
if isinstance(prev_sample, torch.Tensor | float) and not return_dict:
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
@@ -25,7 +25,6 @@ from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.vaes.common import (DiagonalGaussianDistribution,
|
||||
ParallelTiledVAE)
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
@@ -93,8 +92,6 @@ class WanCausalConv3d(nn.Conv3d):
|
||||
x = torch.cat([cache_x, x], dim=2)
|
||||
padding[4] -= cache_x.shape[2]
|
||||
x = F.pad(x, padding)
|
||||
x = x.to(self.weight.dtype) if current_platform.is_mps(
|
||||
) else x # casting needed for mps since amp isn't supported
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.layers.lora.linear import (BaseLayerWithLoRA, get_lora_layer,
|
||||
replace_submodule)
|
||||
@@ -34,7 +34,6 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.device = get_local_torch_device()
|
||||
self.exclude_lora_layers = self.modules[
|
||||
"transformer"].config.arch_config.exclude_lora_layers
|
||||
self.device = get_local_torch_device()
|
||||
|
||||
@@ -3,15 +3,11 @@
|
||||
Decoding stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import weakref
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
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.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -30,9 +26,8 @@ class DecodingStage(PipelineStage):
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
|
||||
def __init__(self, vae, pipeline=None) -> None:
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
@@ -66,15 +61,6 @@ class DecodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
loader = VAELoader()
|
||||
self.vae = loader.load(fastvideo_args.model_paths["vae"],
|
||||
fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
@@ -134,12 +120,4 @@ class DecodingStage(PipelineStage):
|
||||
|
||||
self.vae.to("cpu")
|
||||
|
||||
if torch.backends.mps.is_available():
|
||||
del self.vae
|
||||
if pipeline is not None and "vae" in pipeline.modules:
|
||||
del pipeline.modules["vae"]
|
||||
gc.collect()
|
||||
torch.mps.empty_cache()
|
||||
fastvideo_args.model_loaded["vae"] = False
|
||||
|
||||
return batch
|
||||
|
||||
@@ -3,9 +3,7 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import inspect
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
@@ -23,7 +21,6 @@ from fastvideo.v1.distributed.communication_op import (
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
@@ -56,11 +53,10 @@ class DenoisingStage(PipelineStage):
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None) -> None:
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
@@ -87,15 +83,6 @@ class DenoisingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
@@ -314,19 +301,6 @@ class DenoisingStage(PipelineStage):
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
# deallocate transformer if on mps
|
||||
if torch.backends.mps.is_available():
|
||||
logger.info("Memory before deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
del self.transformer
|
||||
if pipeline is not None and "transformer" in pipeline.modules:
|
||||
del pipeline.modules["transformer"]
|
||||
gc.collect()
|
||||
torch.mps.empty_cache()
|
||||
fastvideo_args.model_loaded["transformer"] = False
|
||||
logger.info("Memory after deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
|
||||
return batch
|
||||
|
||||
def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]:
|
||||
|
||||
@@ -8,13 +8,12 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
logger = (__name__)
|
||||
|
||||
|
||||
class TextEncodingStage(PipelineStage):
|
||||
@@ -64,7 +63,7 @@ class TextEncodingStage(PipelineStage):
|
||||
fastvideo_args.pipeline_config.postprocess_text_funcs,
|
||||
strict=True):
|
||||
|
||||
assert isinstance(batch.prompt, str | list)
|
||||
assert isinstance(batch.prompt, (str | list))
|
||||
if isinstance(batch.prompt, str):
|
||||
batch.prompt = [batch.prompt]
|
||||
texts = []
|
||||
|
||||
@@ -28,12 +28,12 @@ class StageValidators:
|
||||
@staticmethod
|
||||
def positive_float(value: Any) -> bool:
|
||||
"""Check if value is a positive float."""
|
||||
return isinstance(value, int | float) and value > 0
|
||||
return isinstance(value, (int | float)) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def non_negative_float(value: Any) -> bool:
|
||||
"""Check if value is a non-negative float."""
|
||||
return isinstance(value, int | float) and value >= 0
|
||||
return isinstance(value, (int | float)) and value >= 0
|
||||
|
||||
@staticmethod
|
||||
def divisible_by(value: Any, divisor: int) -> bool:
|
||||
|
||||
@@ -61,12 +61,10 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
|
||||
@@ -47,57 +47,19 @@ def cuda_platform_plugin() -> str | None:
|
||||
return "fastvideo.v1.platforms.cuda.CudaPlatform" if is_cuda else None
|
||||
|
||||
|
||||
def mps_platform_plugin() -> str | None:
|
||||
"""Detect if MPS (Metal Performance Shaders) is available on macOS."""
|
||||
is_mps = False
|
||||
|
||||
try:
|
||||
import torch
|
||||
if torch.backends.mps.is_available():
|
||||
is_mps = True
|
||||
logger.info("MPS (Metal Performance Shaders) is available")
|
||||
else:
|
||||
logger.info("MPS is not available")
|
||||
except Exception as e:
|
||||
logger.info("MPS detection failed: %s", e)
|
||||
|
||||
return "fastvideo.v1.platforms.mps.MpsPlatform" if is_mps else None
|
||||
|
||||
|
||||
def cpu_platform_plugin() -> str | None:
|
||||
"""Detect if CPU platform should be used."""
|
||||
# CPU is always available as a fallback
|
||||
return "fastvideo.v1.platforms.cpu.CpuPlatform"
|
||||
|
||||
|
||||
builtin_platform_plugins = {
|
||||
'cuda': cuda_platform_plugin,
|
||||
'mps': mps_platform_plugin,
|
||||
'cpu': cpu_platform_plugin,
|
||||
}
|
||||
|
||||
|
||||
def resolve_current_platform_cls_qualname() -> str:
|
||||
# TODO(will): if we need to support other platforms, we should consider if
|
||||
# vLLM's plugin architecture is suitable for our needs.
|
||||
|
||||
# Try MPS first on macOS
|
||||
platform_cls_qualname = mps_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to CUDA
|
||||
platform_cls_qualname = cuda_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to CPU as last resort
|
||||
platform_cls_qualname = cpu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
raise RuntimeError("No platform plugin found. Please check your "
|
||||
"installation.")
|
||||
platform_cls_qualname = builtin_platform_plugins['cuda']()
|
||||
if platform_cls_qualname is None:
|
||||
raise RuntimeError("No platform plugin found. Please check your "
|
||||
"installation.")
|
||||
return platform_cls_qualname
|
||||
|
||||
|
||||
_current_platform = None
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/cpu.py
|
||||
|
||||
import platform
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.platforms.interface import CpuArchEnum, Platform, PlatformEnum
|
||||
|
||||
|
||||
class CpuPlatform(Platform):
|
||||
_enum = PlatformEnum.CPU
|
||||
device_name = "CPU"
|
||||
device_type = "cpu"
|
||||
dispatch_key = "CPU"
|
||||
simple_compile_backend = "inductor"
|
||||
supported_quantization = []
|
||||
|
||||
@classmethod
|
||||
def get_cpu_architecture(cls) -> CpuArchEnum:
|
||||
"""Get the CPU architecture."""
|
||||
machine = platform.machine().lower()
|
||||
if machine in ("x86_64", "amd64", "i386", "i686"):
|
||||
return CpuArchEnum.X86
|
||||
elif machine in ("arm64", "aarch64"):
|
||||
return CpuArchEnum.ARM
|
||||
else:
|
||||
return CpuArchEnum.UNSPECIFIED
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return platform.processor()
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
return platform.machine()
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
# This is a rough estimate for CPU memory
|
||||
# In practice, you might want to use psutil or similar
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls,
|
||||
device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
# For CPU, we can't easily get memory usage without additional libraries
|
||||
return 0.0
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "fastvideo.v1.distributed.device_communicators.cpu_communicator.CpuCommunicator"
|
||||
@@ -27,17 +27,10 @@ class PlatformEnum(enum.Enum):
|
||||
ROCM = enum.auto()
|
||||
TPU = enum.auto()
|
||||
CPU = enum.auto()
|
||||
MPS = enum.auto()
|
||||
OOT = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
|
||||
|
||||
class CpuArchEnum(enum.Enum):
|
||||
X86 = enum.auto()
|
||||
ARM = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
|
||||
|
||||
class DeviceCapability(NamedTuple):
|
||||
major: int
|
||||
minor: int
|
||||
@@ -94,9 +87,6 @@ class Platform:
|
||||
# TODO(will): ROCM will be supported in the future here
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
def is_mps(self) -> bool:
|
||||
return self._enum == PlatformEnum.MPS
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
@@ -219,11 +209,6 @@ class Platform:
|
||||
"""
|
||||
return "fastvideo.v1.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase" # noqa
|
||||
|
||||
@classmethod
|
||||
def get_cpu_architecture(cls) -> CpuArchEnum:
|
||||
"""Get the CPU architecture of the current platform."""
|
||||
return CpuArchEnum.UNSPECIFIED
|
||||
|
||||
|
||||
class UnspecifiedPlatform(Platform):
|
||||
_enum = PlatformEnum.UNSPECIFIED
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.platforms.interface import (DeviceCapability, Platform,
|
||||
PlatformEnum)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MpsPlatform(Platform):
|
||||
_enum = PlatformEnum.MPS
|
||||
device_name: str = "mps"
|
||||
device_type: str = "mps"
|
||||
dispatch_key: str = "MPS"
|
||||
device_control_env_var: str = "MPS_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls,
|
||||
device_id: int = 0) -> DeviceCapability | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable MPS "
|
||||
"graph. Since, enforce-eager is enabled, async output "
|
||||
"processor cannot be used")
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls,
|
||||
device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
return 0.0
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible
|
||||
logger.info("Using Torch SDPA backend for MPS.")
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
# Use base communicator for MPS
|
||||
return "fastvideo.v1.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase"
|
||||
|
||||
@classmethod
|
||||
def seed_everything(cls, seed: int | None = None) -> None:
|
||||
"""Set the seed for MPS device."""
|
||||
if seed is not None:
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
# MPS doesn't have manual_seed_all like CUDA
|
||||
# The manual_seed above should be sufficient
|
||||
+7
-20
@@ -80,17 +80,16 @@ prev_set_stream = torch.cuda.set_stream
|
||||
_current_stream = None
|
||||
|
||||
|
||||
def _patched_set_stream(stream: torch.cuda.Stream | None) -> None:
|
||||
def _patched_set_stream(stream: torch.cuda.Stream) -> None:
|
||||
global _current_stream
|
||||
_current_stream = stream
|
||||
if stream is not None:
|
||||
prev_set_stream(stream)
|
||||
prev_set_stream(stream)
|
||||
|
||||
|
||||
torch.cuda.set_stream = _patched_set_stream
|
||||
|
||||
|
||||
def current_stream() -> torch.cuda.Stream | None:
|
||||
def current_stream() -> torch.cuda.Stream:
|
||||
"""
|
||||
replace `torch.cuda.current_stream()` with `fastvideo.v1.utils.current_stream()`.
|
||||
it turns out that `torch.cuda.current_stream()` is quite expensive,
|
||||
@@ -102,11 +101,6 @@ def current_stream() -> torch.cuda.Stream | None:
|
||||
from C/C++ code.
|
||||
"""
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
# For non-CUDA platforms, return None
|
||||
if not current_platform.is_cuda_alike():
|
||||
return None
|
||||
|
||||
global _current_stream
|
||||
if _current_stream is None:
|
||||
# when this function is called before any stream is set,
|
||||
@@ -652,21 +646,14 @@ def shallow_asdict(obj) -> dict[str, Any]:
|
||||
return {f.name: getattr(obj, f.name) for f in fields(obj)}
|
||||
|
||||
|
||||
# TODO: validate that this is fine
|
||||
def kill_itself_when_parent_died() -> None:
|
||||
# if sys.platform == "linux":
|
||||
# sigkill this process when parent worker manager dies
|
||||
PR_SET_PDEATHSIG = 1
|
||||
import platform
|
||||
if platform.system() == "Linux":
|
||||
libc = ctypes.CDLL("libc.so.6")
|
||||
libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL)
|
||||
# elif platform.system() == "Darwin":
|
||||
# libc = ctypes.CDLL("libc.dylib")
|
||||
libc = ctypes.CDLL("libc.so.6")
|
||||
libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL)
|
||||
# else:
|
||||
# logger.warning("kill_itself_when_parent_died is only supported in linux.")
|
||||
else:
|
||||
logger.warning(
|
||||
"kill_itself_when_parent_died is only supported in linux.")
|
||||
|
||||
|
||||
def get_exception_traceback() -> str:
|
||||
@@ -799,7 +786,7 @@ def dict_to_3d_list(
|
||||
t, l, h = map(int, key.split("_")) # noqa: E741
|
||||
if 0 <= t < max_timesteps_idx and 0 <= l < max_layer_idx and 0 <= h < max_head_idx:
|
||||
result[t][l][h] = value
|
||||
# else: silently ignore any key that doesn't fit
|
||||
# else: silently ignore any key that doesn’t fit
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@@ -14,11 +14,9 @@ import torch
|
||||
from fastvideo.v1.distributed import (
|
||||
cleanup_dist_env_and_memory,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ForwardBatch, build_pipeline
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from fastvideo.v1.utils import (get_exception_traceback,
|
||||
kill_itself_when_parent_died)
|
||||
|
||||
@@ -66,17 +64,11 @@ class Worker:
|
||||
|
||||
# This env var set by Ray causes exceptions with graph building.
|
||||
os.environ.pop("NCCL_ASYNC_ERROR_HANDLING", None)
|
||||
|
||||
# Platform-agnostic device initialization
|
||||
self.device = get_local_torch_device()
|
||||
self.device = torch.device(f"cuda:{self.local_rank}")
|
||||
torch.cuda.set_device(self.device)
|
||||
|
||||
# _check_if_gpu_supports_dtype(self.model_config.dtype)
|
||||
if current_platform.is_cuda_alike():
|
||||
torch.cuda.empty_cache()
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
|
||||
else:
|
||||
# For MPS, we can't get memory info the same way
|
||||
self.init_gpu_memory = 0
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = str(self.master_port)
|
||||
|
||||
+2
-2
@@ -19,7 +19,7 @@ dependencies = [
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
|
||||
"timm==1.0.11", "peft>=0.15.0", "diffusers>=0.33.1",
|
||||
"timm==1.0.11", "peft>=0.15.0", "diffusers>=0.33.1", "bitsandbytes",
|
||||
"torch==2.7.1", "torchvision",
|
||||
|
||||
# Acceleration & Optimization
|
||||
@@ -27,7 +27,7 @@ dependencies = [
|
||||
|
||||
# Computer Vision & Image Processing
|
||||
"opencv-python==4.10.0.84", "pillow>=10.3.0", "imageio==2.36.0",
|
||||
"imageio-ffmpeg==0.5.1", "einops",
|
||||
"imageio-ffmpeg==0.5.1", "decord==0.6.0", "einops",
|
||||
|
||||
# Experiment Tracking & Logging
|
||||
"wandb>=0.19.11", "loguru", "test-tube==0.7.5",
|
||||
|
||||
Reference in New Issue
Block a user