Compare commits

...
Author SHA1 Message Date
William Lin 3def003890 Revert "video gen working on apple silicon (addressed issues from prior pr) (…"
This reverts commit b79d1fc15b.
2025-07-15 22:04:34 -07:00
33 changed files with 132 additions and 749 deletions
-40
View File
@@ -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()
-39
View File
@@ -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,
+39 -82
View File
@@ -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
+6 -40
View File
@@ -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
+1 -3
View File
@@ -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)
+11 -22
View File
@@ -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)
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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:
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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:
+2 -9
View File
@@ -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
+2 -4
View File
@@ -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()
+12 -27
View File
@@ -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)
+17 -29
View File
@@ -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)
+6 -4
View File
@@ -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)
-3
View File
@@ -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)
+1 -2
View File
@@ -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()
+1 -23
View File
@@ -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
+1 -27
View File
@@ -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 = []
+2 -2
View File
@@ -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:
+2 -4
View File
@@ -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
+5 -43
View File
@@ -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
-57
View File
@@ -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"
-15
View File
@@ -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
-76
View File
@@ -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
View File
@@ -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
+3 -11
View File
@@ -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
View File
@@ -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",